diff --git a/docs/changelog/index.md b/docs/changelog/index.md index 49cdffe0..a0edcd39 100644 --- a/docs/changelog/index.md +++ b/docs/changelog/index.md @@ -18,6 +18,9 @@ sidebar: false ### 变更 +- 访问日志列表查询将分页与计数下推到数据库执行,避免百万级数据全量加载到内存。 +- 访问日志 `total_ip` 统计改为 SQL `UNION` + `COUNT(*)` 下推执行,分片计数与分页查询并行化。 +- 访问日志折叠视图、IP 汇总与趋势改为 SQL `GROUP BY` 聚合;过滤条件改为 `node_id` 精确匹配及其他字段前缀匹配以利用索引。 - 标准化 Server Go 目录结构,引入 `cmd/server`、`openflare-server/internal` 与根级 `pkg` 分层,并拆分原 `utils` 公共能力包。 ## [v2.3.3] - 2026-06-06 diff --git a/openflare-server/.gitignore b/openflare-server/.gitignore index 82f0c3ac..0d28eed5 100644 --- a/openflare-server/.gitignore +++ b/openflare-server/.gitignore @@ -1 +1,2 @@ /data/ +/postgres-data/ diff --git a/openflare-server/cmd/accesslogbench/main.go b/openflare-server/cmd/accesslogbench/main.go new file mode 100644 index 00000000..8aa7b758 --- /dev/null +++ b/openflare-server/cmd/accesslogbench/main.go @@ -0,0 +1,562 @@ +package main + +import ( + "flag" + "fmt" + "log" + "os" + "path/filepath" + "runtime" + "sort" + "strings" + "time" + + "github.com/rain-kl/openflare/openflare-server/internal/common" + "github.com/rain-kl/openflare/openflare-server/internal/model" +) + +const retentionDays = 90 + +func main() { + dsn := flag.String("dsn", envOr("DSN", "postgres://openflare:replace-with-strong-password@192.168.107.2:5432/openflare?sslmode=disable"), "PostgreSQL DSN") + records := flag.Int("records", 1_000_000, "number of access log rows to seed (0 = skip seeding)") + seedNode := flag.String("node-id", "bench-node", "node_id used for seeded rows") + iterations := flag.Int("iterations", 10, "benchmark iterations per scenario") + warmup := flag.Int("warmup", 2, "warmup iterations per scenario") + page := flag.Int("page", 0, "page index for list benchmark") + pageSize := flag.Int("page-size", 20, "page size for list benchmark") + reset := flag.Bool("reset", false, "truncate node_access_logs shard tables before seeding") + legacyIterations := flag.Int("legacy-iterations", 1, "iterations for legacy full-scan scenario (can be very slow)") + skipLegacy := flag.Bool("skip-legacy", false, "skip legacy full-scan scenario") + verifyOnly := flag.Bool("verify-only", false, "verify query correctness and exit") + flag.Parse() + + common.SQLDSN = *dsn + common.SQLitePath = filepath.Join(os.TempDir(), "openflare-accesslogbench-missing.sqlite") + if err := initBenchDB(*dsn); err != nil { + log.Fatalf("init db: %v", err) + } + + fmt.Println("=== OpenFlare node_access_logs benchmark (PostgreSQL) ===") + fmt.Printf("DSN: %s\n", redactDSN(*dsn)) + fmt.Printf("GOMAXPROCS=%d\n", runtime.GOMAXPROCS(0)) + + if *reset { + if err := truncateAccessLogs(); err != nil { + log.Fatalf("truncate access logs: %v", err) + } + fmt.Println("truncated node_access_logs shard tables") + } + + if *records > 0 { + existing, err := countAllAccessLogs() + if err != nil { + log.Fatalf("count existing rows: %v", err) + } + if existing >= int64(*records) { + fmt.Printf("existing rows=%d >= target=%d, skip seeding\n", existing, *records) + } else { + missing := *records - int(existing) + fmt.Printf("seeding %d rows (existing=%d)...\n", missing, existing) + if err := seedAccessLogs(*seedNode, missing); err != nil { + log.Fatalf("seed access logs: %v", err) + } + } + } + + total, err := countAllAccessLogs() + if err != nil { + log.Fatalf("count rows after seed: %v", err) + } + fmt.Printf("total rows across shards: %d\n\n", total) + + since := time.Now().UTC().Add(-retentionDays * 24 * time.Hour) + if err := verifyQueryCorrectness(since, total); err != nil { + log.Fatalf("correctness verification failed: %v", err) + } + fmt.Println("correctness verification: PASS") + if *verifyOnly { + return + } + + query := model.NodeAccessLogQuery{ + Since: since, + Page: *page, + PageSize: *pageSize, + SortBy: "logged_at", + SortOrder: "desc", + } + + scenarios := []scenario{ + { + name: "paginated_list", + run: func() error { + _, err := model.ListNodeAccessLogs(query) + return err + }, + }, + { + name: "sql_count", + run: func() error { + _, _, err := model.CountNodeAccessLogs(query) + return err + }, + }, + { + name: "api_list+count", + run: func() error { + if _, err := model.ListNodeAccessLogs(query); err != nil { + return err + } + _, _, err := model.CountNodeAccessLogs(query) + return err + }, + }, + { + name: "fullscan_list_legacy", + run: func() error { + _, err := fullScanList(query) + return err + }, + }, + } + + if err := printExplainPlans(since); err != nil { + log.Printf("warn: explain analyze failed: %v", err) + } + + for _, item := range scenarios { + if item.name == "fullscan_list_legacy" && *skipLegacy { + fmt.Println("[fullscan_list_legacy] skipped") + continue + } + iterationCount := *iterations + if item.name == "fullscan_list_legacy" { + iterationCount = *legacyIterations + } + result, err := runScenario(item, *warmup, iterationCount) + if err != nil { + log.Fatalf("scenario %s failed: %v", item.name, err) + } + printResult(result) + } +} + +type scenario struct { + name string + run func() error +} + +type benchResult struct { + name string + iterations int + latencies []time.Duration + allocBytes []uint64 + heapInUse []uint64 + maxHeapInUse uint64 +} + +func runScenario(item scenario, warmup int, iterations int) (*benchResult, error) { + for range warmup { + if err := item.run(); err != nil { + return nil, err + } + } + + result := &benchResult{ + name: item.name, + iterations: iterations, + } + var peakHeap uint64 + + for range iterations { + runtime.GC() + var before, after runtime.MemStats + runtime.ReadMemStats(&before) + + start := time.Now() + if err := item.run(); err != nil { + return nil, err + } + elapsed := time.Since(start) + + runtime.ReadMemStats(&after) + result.latencies = append(result.latencies, elapsed) + result.allocBytes = append(result.allocBytes, after.TotalAlloc-before.TotalAlloc) + result.heapInUse = append(result.heapInUse, after.HeapInuse) + if after.HeapInuse > peakHeap { + peakHeap = after.HeapInuse + } + } + result.maxHeapInUse = peakHeap + return result, nil +} + +func printResult(result *benchResult) { + latencies := append([]time.Duration(nil), result.latencies...) + sort.Slice(latencies, func(i, j int) bool { return latencies[i] < latencies[j] }) + + var allocSum uint64 + var heapSum uint64 + for index := range result.allocBytes { + allocSum += result.allocBytes[index] + heapSum += result.heapInUse[index] + } + + fmt.Printf("[%s] iterations=%d\n", result.name, result.iterations) + fmt.Printf(" latency: min=%s avg=%s p50=%s p95=%s max=%s\n", + minDuration(latencies), + avgDuration(latencies), + percentile(latencies, 50), + percentile(latencies, 95), + maxDuration(latencies), + ) + fmt.Printf(" alloc/op: avg=%s peak_heap=%s\n", + humanBytes(allocSum/uint64(len(result.allocBytes))), + humanBytes(result.maxHeapInUse), + ) + fmt.Printf(" heap_inuse/op: avg=%s\n\n", humanBytes(heapSum/uint64(len(result.heapInUse)))) +} + +func seedAccessLogs(nodeID string, total int) error { + started := time.Now() + perShard := total / 10 + remainder := total % 10 + seeded := 0 + now := time.Now().UTC() + + for shard := range 10 { + rows := perShard + if shard < remainder { + rows++ + } + if rows == 0 { + continue + } + table := fmt.Sprintf("node_access_logs_%02d", shard) + sql := fmt.Sprintf(` +INSERT INTO %s (id, node_id, logged_at, remote_addr, region, host, path, status_code, created_at) +SELECT + (gs * 10 + %d)::bigint AS id, + $1 AS node_id, + $2::timestamptz - (gs || ' minutes')::interval AS logged_at, + ('203.0.' || ((gs / 256) %% 256)::text || '.' || (gs %% 256)::text) AS remote_addr, + 'Benchland' AS region, + ('host-' || (gs %% 200)::text || '.example.com') AS host, + ('/api/v1/resource/' || gs::text) AS path, + (200 + (gs %% 4))::bigint AS status_code, + $2::timestamptz AS created_at +FROM generate_series(0, $3 - 1) AS gs +`, table, shard) + if err := model.DB.Exec(sql, nodeID, now, rows).Error; err != nil { + return err + } + seeded += rows + elapsed := time.Since(started) + rate := float64(seeded) / elapsed.Seconds() + fmt.Printf(" seeded %d/%d rows (%.0f rows/s, elapsed %s)\n", seeded, total, rate, elapsed.Round(time.Millisecond)) + } + return nil +} + +func truncateAccessLogs() error { + for _, table := range observabilityShardTables() { + if err := model.DB.Exec("TRUNCATE TABLE " + table).Error; err != nil { + return err + } + } + return nil +} + +func countAllAccessLogs() (int64, error) { + var total int64 + for _, table := range observabilityShardTables() { + var count int64 + if err := model.DB.Table(table).Count(&count).Error; err != nil { + return 0, err + } + total += count + } + return total, nil +} + +func observabilityShardTables() []string { + tables := make([]string, 0, 10) + for index := range 10 { + tables = append(tables, fmt.Sprintf("node_access_logs_%02d", index)) + } + return tables +} + +func verifyQueryCorrectness(since time.Time, tableTotal int64) error { + fmt.Println("=== correctness verification ===") + baseQuery := model.NodeAccessLogQuery{ + Since: since, + SortBy: "logged_at", + SortOrder: "desc", + } + + reference, err := fullScanList(baseQuery) + if err != nil { + return fmt.Errorf("reference full scan failed: %w", err) + } + if int64(len(reference)) != tableTotal { + return fmt.Errorf("reference row count %d != table total %d", len(reference), tableTotal) + } + fmt.Printf(" reference rows in retention window: %d\n", len(reference)) + + totalRecords, totalIPs, err := model.CountNodeAccessLogs(baseQuery) + if err != nil { + return fmt.Errorf("CountNodeAccessLogs failed: %w", err) + } + if totalRecords != int64(len(reference)) { + return fmt.Errorf("total_records=%d want reference=%d", totalRecords, len(reference)) + } + referenceIPs := countUniqueIPs(reference) + if totalIPs != referenceIPs { + return fmt.Errorf("total_ip=%d want reference=%d", totalIPs, referenceIPs) + } + fmt.Printf(" count totals: total_record=%d total_ip=%d\n", totalRecords, totalIPs) + + pages := []struct { + page int + pageSize int + }{ + {0, 20}, + {1, 20}, + {49, 20}, + {100, 50}, + } + for _, item := range pages { + query := baseQuery + query.Page = item.page + query.PageSize = item.pageSize + got, err := model.ListNodeAccessLogs(query) + if err != nil { + return fmt.Errorf("ListNodeAccessLogs page=%d size=%d failed: %w", item.page, item.pageSize, err) + } + start := item.page * item.pageSize + end := start + item.pageSize + if start >= len(reference) { + if len(got) != 0 { + return fmt.Errorf("page=%d size=%d expected empty got %d", item.page, item.pageSize, len(got)) + } + continue + } + if end > len(reference) { + end = len(reference) + } + want := reference[start:end] + if !accessLogsEqual(got, want) { + return fmt.Errorf("page=%d size=%d content mismatch (got %d want %d rows)", item.page, item.pageSize, len(got), len(want)) + } + fmt.Printf(" page=%d size=%d: %d rows match reference\n", item.page, item.pageSize, len(got)) + } + + filtered := model.NodeAccessLogQuery{ + NodeID: "bench-node", + Since: since, + SortBy: "status_code", + SortOrder: "asc", + Page: 2, + PageSize: 15, + } + filteredReference, err := fullScanList(filtered) + if err != nil { + return fmt.Errorf("filtered reference failed: %w", err) + } + filteredRows, err := model.ListNodeAccessLogs(filtered) + if err != nil { + return fmt.Errorf("filtered ListNodeAccessLogs failed: %w", err) + } + start := filtered.Page * filtered.PageSize + end := start + filtered.PageSize + if end > len(filteredReference) { + end = len(filteredReference) + } + if start >= len(filteredReference) { + start = len(filteredReference) + } + if !accessLogsEqual(filteredRows, filteredReference[start:end]) { + return fmt.Errorf("filtered page content mismatch") + } + filteredTotal, filteredIPs, err := model.CountNodeAccessLogs(filtered) + if err != nil { + return fmt.Errorf("filtered CountNodeAccessLogs failed: %w", err) + } + if filteredTotal != int64(len(filteredReference)) { + return fmt.Errorf("filtered total_records=%d want %d", filteredTotal, len(filteredReference)) + } + if filteredIPs != countUniqueIPs(filteredReference) { + return fmt.Errorf("filtered total_ip=%d want %d", filteredIPs, countUniqueIPs(filteredReference)) + } + fmt.Printf(" filtered node_id=bench-node page=2 size=15: match (total_record=%d total_ip=%d)\n", filteredTotal, filteredIPs) + return nil +} + +func countUniqueIPs(logs []*model.NodeAccessLog) int64 { + ips := make(map[string]struct{}) + for _, item := range logs { + if item == nil { + continue + } + if trimmed := strings.TrimSpace(item.RemoteAddr); trimmed != "" { + ips[trimmed] = struct{}{} + } + } + return int64(len(ips)) +} + +func accessLogsEqual(left []*model.NodeAccessLog, right []*model.NodeAccessLog) bool { + if len(left) != len(right) { + return false + } + for index := range left { + if left[index] == nil || right[index] == nil { + if left[index] != right[index] { + return false + } + continue + } + if left[index].ID != right[index].ID || + left[index].NodeID != right[index].NodeID || + !left[index].LoggedAt.Equal(right[index].LoggedAt) || + left[index].RemoteAddr != right[index].RemoteAddr || + left[index].Host != right[index].Host || + left[index].Path != right[index].Path || + left[index].StatusCode != right[index].StatusCode { + return false + } + } + return true +} + +func printExplainPlans(since time.Time) error { + fmt.Println("=== PostgreSQL EXPLAIN (one shard sample: node_access_logs_00) ===") + queries := []struct { + name string + sql string + }{ + { + name: "paginated_list", + sql: "EXPLAIN (ANALYZE, BUFFERS) SELECT * FROM node_access_logs_00 WHERE logged_at >= $1 ORDER BY logged_at DESC, id DESC LIMIT 20", + }, + { + name: "count_rows", + sql: "EXPLAIN (ANALYZE, BUFFERS) SELECT COUNT(*) FROM node_access_logs_00 WHERE logged_at >= $1", + }, + { + name: "distinct_ip_union_all", + sql: `EXPLAIN (ANALYZE, BUFFERS) SELECT COUNT(*) FROM ( +SELECT remote_addr FROM ( +SELECT TRIM(remote_addr) AS remote_addr FROM node_access_logs_00 WHERE logged_at >= $1 AND remote_addr <> '' +UNION ALL +SELECT TRIM(remote_addr) AS remote_addr FROM node_access_logs_01 WHERE logged_at >= $1 AND remote_addr <> '' +) AS all_ips GROUP BY remote_addr +) AS ips`, + }, + } + for _, item := range queries { + rows, err := model.DB.Raw(item.sql, since).Rows() + if err != nil { + return err + } + fmt.Printf("-- %s\n", item.name) + for rows.Next() { + var line string + if err := rows.Scan(&line); err != nil { + rows.Close() + return err + } + fmt.Println(line) + } + rows.Close() + fmt.Println() + } + return nil +} + +func fullScanList(query model.NodeAccessLogQuery) ([]*model.NodeAccessLog, error) { + legacy := query + legacy.PageSize = 0 + return model.ListNodeAccessLogs(legacy) +} + +func initBenchDB(dsn string) error { + return model.InitBenchmarkDB(dsn) +} + +func envOr(key string, fallback string) string { + if value := strings.TrimSpace(os.Getenv(key)); value != "" { + return value + } + return fallback +} + +func redactDSN(dsn string) string { + if at := strings.Index(dsn, "@"); at > 0 { + schemeEnd := strings.Index(dsn, "://") + if schemeEnd >= 0 { + return dsn[:schemeEnd+3] + "***@" + dsn[at+1:] + } + } + return dsn +} + +func minDuration(values []time.Duration) time.Duration { + if len(values) == 0 { + return 0 + } + return values[0] +} + +func maxDuration(values []time.Duration) time.Duration { + if len(values) == 0 { + return 0 + } + return values[len(values)-1] +} + +func avgDuration(values []time.Duration) time.Duration { + if len(values) == 0 { + return 0 + } + var sum time.Duration + for _, value := range values { + sum += value + } + return sum / time.Duration(len(values)) +} + +func percentile(values []time.Duration, p int) time.Duration { + if len(values) == 0 { + return 0 + } + if p <= 0 { + return values[0] + } + if p >= 100 { + return values[len(values)-1] + } + index := (len(values)*p + 99) / 100 + if index <= 0 { + index = 1 + } + if index > len(values) { + index = len(values) + } + return values[index-1] +} + +func humanBytes(value uint64) string { + const unit = 1024 + if value < unit { + return fmt.Sprintf("%d B", value) + } + div, exp := uint64(unit), 0 + for n := value / unit; n >= unit; n /= unit { + div *= unit + exp++ + } + return fmt.Sprintf("%.1f %ciB", float64(value)/float64(div), "KMGTPE"[exp]) +} diff --git a/openflare-server/internal/model/main.go b/openflare-server/internal/model/main.go index 13c34a7c..02c18af7 100644 --- a/openflare-server/internal/model/main.go +++ b/openflare-server/internal/model/main.go @@ -336,6 +336,24 @@ func resetPostgresSequence(db *gorm.DB, tableName string) error { return db.Exec(sql).Error } +func InitBenchmarkDB(dsn string) error { + db, err := gorm.Open(postgres.Open(dsn), &gorm.Config{}) + if err != nil { + return err + } + sqlDB, err := db.DB() + if err != nil { + return err + } + sqlDB.SetMaxOpenConns(20) + sqlDB.SetMaxIdleConns(10) + DB = db + if err = registerSharding(db, "postgres"); err != nil { + return err + } + return ensureDatabaseSchemaUpToDate(db, "postgres") +} + func InitDB() (err error) { db, backend, err := openDatabase() if err != nil { diff --git a/openflare-server/internal/model/node_access_log.go b/openflare-server/internal/model/node_access_log.go index 6b41826d..4452bea7 100644 --- a/openflare-server/internal/model/node_access_log.go +++ b/openflare-server/internal/model/node_access_log.go @@ -1,8 +1,10 @@ package model import ( + "fmt" "sort" "strings" + "sync" "time" "gorm.io/gorm" @@ -119,15 +121,10 @@ func (log *NodeAccessLog) BeforeCreate(*gorm.DB) error { } func ListNodeAccessLogs(query NodeAccessLogQuery) (logs []*NodeAccessLog, err error) { - all, err := listNodeAccessLogsAcrossShards(query) - if err != nil { - return nil, err + if query.PageSize > 0 { + return listNodeAccessLogsPaginatedAcrossShards(query) } - start, end := paginateBounds(len(all), query.Page, query.PageSize) - if start >= len(all) { - return []*NodeAccessLog{}, nil - } - return all[start:end], nil + return listNodeAccessLogsAcrossShards(query) } func ListNodeAccessLogsForWAFIPGroup(query NodeAccessLogQuery) ([]*NodeAccessLog, error) { @@ -135,21 +132,27 @@ func ListNodeAccessLogsForWAFIPGroup(query NodeAccessLogQuery) ([]*NodeAccessLog } func CountNodeAccessLogs(query NodeAccessLogQuery) (totalRecords int64, totalIPs int64, err error) { - all, err := listNodeAccessLogsAcrossShards(query) - if err != nil { - return 0, 0, err + db := normalizeShardedDB(DB) + var countErr error + var distinctErr error + var wg sync.WaitGroup + wg.Add(2) + go func() { + defer wg.Done() + totalRecords, countErr = countNodeAccessLogRecordsAcrossShards(db, query) + }() + go func() { + defer wg.Done() + totalIPs, distinctErr = countDistinctNodeAccessLogIPsAcrossShards(db, query) + }() + wg.Wait() + if countErr != nil { + return 0, 0, countErr } - ips := make(map[string]struct{}, len(all)) - for _, item := range all { - if item == nil { - continue - } - trimmed := strings.TrimSpace(item.RemoteAddr) - if trimmed != "" { - ips[trimmed] = struct{}{} - } + if distinctErr != nil { + return 0, 0, distinctErr } - return int64(len(all)), int64(len(ips)), nil + return totalRecords, totalIPs, nil } func ListNodeAccessLogRegionCounts(nodeID string, since time.Time, limit int) (items []*NodeAccessLogRegionCount, err error) { @@ -251,38 +254,7 @@ func CountNodeAccessLogIPSummaries(query NodeAccessLogIPSummaryQuery) (total int } func ListNodeAccessLogIPTrend(query NodeAccessLogIPTrendQuery) (items []*NodeAccessLogTrendPointRow, err error) { - logs, err := listNodeAccessLogsAcrossShards(NodeAccessLogQuery{ - NodeID: query.NodeID, - RemoteAddr: query.RemoteAddr, - Host: query.Host, - Since: query.Since, - }) - if err != nil { - return nil, err - } - remoteAddr := strings.TrimSpace(query.RemoteAddr) - if remoteAddr == "" { - return []*NodeAccessLogTrendPointRow{}, nil - } - buckets := make(map[int64]int64) - for _, item := range logs { - if item == nil || strings.TrimSpace(item.RemoteAddr) != remoteAddr { - continue - } - bucketEpoch := bucketEpochForTime(item.LoggedAt, query.BucketMinutes) - buckets[bucketEpoch]++ - } - items = make([]*NodeAccessLogTrendPointRow, 0, len(buckets)) - for bucketEpoch, requestCount := range buckets { - items = append(items, &NodeAccessLogTrendPointRow{ - BucketEpoch: bucketEpoch, - RequestCount: requestCount, - }) - } - sort.Slice(items, func(i int, j int) bool { - return items[i].BucketEpoch < items[j].BucketEpoch - }) - return items, nil + return queryIPTrendRows(query) } func DeleteNodeAccessLogsBefore(before time.Time) (deleted int64, err error) { @@ -329,26 +301,98 @@ func DeleteNodeAccessLogsByNodeBefore(db *gorm.DB, nodeID string, before time.Ti }) } -func applyNodeAccessLogFilters(db *gorm.DB, query NodeAccessLogQuery) *gorm.DB { +func buildNodeAccessLogFilterClause(query NodeAccessLogQuery) (string, []any) { + parts := make([]string, 0, 6) + args := make([]any, 0, 6) if trimmed := strings.TrimSpace(query.NodeID); trimmed != "" { - db = db.Where("node_id LIKE ?", "%"+trimmed+"%") + parts = append(parts, "node_id = ?") + args = append(args, trimmed) } if trimmed := strings.TrimSpace(query.RemoteAddr); trimmed != "" { - db = db.Where("remote_addr LIKE ?", "%"+trimmed+"%") + parts = append(parts, "remote_addr LIKE ?") + args = append(args, trimmed+"%") } if trimmed := strings.TrimSpace(query.Host); trimmed != "" { - db = db.Where("host LIKE ?", "%"+trimmed+"%") + parts = append(parts, "host LIKE ?") + args = append(args, trimmed+"%") } if trimmed := strings.TrimSpace(query.Path); trimmed != "" { - db = db.Where("path LIKE ?", "%"+trimmed+"%") + parts = append(parts, "path LIKE ?") + args = append(args, trimmed+"%") } if !query.Since.IsZero() { - db = db.Where("logged_at >= ?", query.Since) + parts = append(parts, "logged_at >= ?") + args = append(args, query.Since) } if !query.Until.IsZero() { - db = db.Where("logged_at < ?", query.Until) + parts = append(parts, "logged_at < ?") + args = append(args, query.Until) } - return db + if len(parts) == 0 { + return "TRUE", nil + } + return strings.Join(parts, " AND "), args +} + +func applyNodeAccessLogFilters(db *gorm.DB, query NodeAccessLogQuery) *gorm.DB { + clause, args := buildNodeAccessLogFilterClause(query) + if clause == "TRUE" { + return db + } + return db.Where(clause, args...) +} + +func countNodeAccessLogRecordsAcrossShards(db *gorm.DB, query NodeAccessLogQuery) (int64, error) { + tables := observabilityShardTables("node_access_logs") + counts := make([]int64, len(tables)) + errs := make([]error, len(tables)) + + var wg sync.WaitGroup + for index, table := range tables { + wg.Add(1) + go func(index int, table string) { + defer wg.Done() + var count int64 + errs[index] = applyNodeAccessLogFilters(db.Table(table), query).Count(&count).Error + counts[index] = count + }(index, table) + } + wg.Wait() + + var total int64 + for index := range tables { + if errs[index] != nil { + return 0, errs[index] + } + total += counts[index] + } + return total, nil +} + +func countDistinctNodeAccessLogIPsAcrossShards(db *gorm.DB, query NodeAccessLogQuery) (int64, error) { + clause, args := buildNodeAccessLogFilterClause(query) + tables := observabilityShardTables("node_access_logs") + unionParts := make([]string, 0, len(tables)) + allArgs := make([]any, 0, len(args)*len(tables)) + for _, table := range tables { + unionParts = append(unionParts, fmt.Sprintf( + "SELECT TRIM(remote_addr) AS remote_addr FROM %s WHERE %s AND remote_addr <> ''", + table, + clause, + )) + allArgs = append(allArgs, args...) + } + sql := fmt.Sprintf(` +SELECT COUNT(*) FROM ( + SELECT remote_addr + FROM (%s) AS all_ips + GROUP BY remote_addr +) AS ips`, strings.Join(unionParts, " UNION ALL ")) + var total int64 + if err := db.Raw(sql, allArgs...).Scan(&total).Error; err != nil { + return 0, err + } + return total, nil } func listNodeAccessLogsAcrossShards(query NodeAccessLogQuery) ([]*NodeAccessLog, error) { @@ -366,188 +410,59 @@ func listNodeAccessLogsAcrossShards(query NodeAccessLogQuery) ([]*NodeAccessLog, return items, nil } -func buildNodeAccessLogBucketRows(query NodeAccessLogBucketQuery) ([]*NodeAccessLogBucketRow, error) { - logs, err := listNodeAccessLogsAcrossShards(NodeAccessLogQuery{ - NodeID: query.NodeID, - RemoteAddr: query.RemoteAddr, - Host: query.Host, - Path: query.Path, - Since: query.Since, - }) - if err != nil { - return nil, err +func listNodeAccessLogsPaginatedAcrossShards(query NodeAccessLogQuery) ([]*NodeAccessLog, error) { + fetchLimit := nodeAccessLogFetchLimit(query.Page, query.PageSize) + orderClause := nodeAccessLogOrderClause(query.SortBy, query.SortOrder) + + items := make([]*NodeAccessLog, 0, fetchLimit*observabilityShardCount) + db := normalizeShardedDB(DB) + for _, table := range observabilityShardTables("node_access_logs") { + var shardRows []*NodeAccessLog + tx := applyNodeAccessLogFilters(db.Table(table), query).Order(orderClause).Limit(fetchLimit) + if err := tx.Find(&shardRows).Error; err != nil { + return nil, err + } + items = append(items, shardRows...) } - type bucketAccumulator struct { - requestCount int64 - uniqueIPs map[string]struct{} - uniqueHosts map[string]struct{} - successCount int64 - clientErrorCount int64 - serverErrorCount int64 + + sortNodeAccessLogs(items, query.SortBy, query.SortOrder) + start, end := paginateBounds(len(items), query.Page, query.PageSize) + if start >= len(items) { + return []*NodeAccessLog{}, nil } - accumulators := make(map[int64]*bucketAccumulator) - for _, item := range logs { - if item == nil { - continue - } - bucketEpoch := bucketEpochForTime(item.LoggedAt, query.FoldMinutes) - accumulator := accumulators[bucketEpoch] - if accumulator == nil { - accumulator = &bucketAccumulator{ - uniqueIPs: make(map[string]struct{}), - uniqueHosts: make(map[string]struct{}), - } - accumulators[bucketEpoch] = accumulator - } - accumulator.requestCount++ - if trimmed := strings.TrimSpace(item.RemoteAddr); trimmed != "" { - accumulator.uniqueIPs[trimmed] = struct{}{} - } - if trimmed := strings.TrimSpace(item.Host); trimmed != "" { - accumulator.uniqueHosts[trimmed] = struct{}{} - } - switch { - case item.StatusCode < 400: - accumulator.successCount++ - case item.StatusCode < 500: - accumulator.clientErrorCount++ - default: - accumulator.serverErrorCount++ - } - } - rows := make([]*NodeAccessLogBucketRow, 0, len(accumulators)) - for bucketEpoch, accumulator := range accumulators { - rows = append(rows, &NodeAccessLogBucketRow{ - BucketEpoch: bucketEpoch, - RequestCount: accumulator.requestCount, - UniqueIPCount: int64(len(accumulator.uniqueIPs)), - UniqueHostCount: int64(len(accumulator.uniqueHosts)), - SuccessCount: accumulator.successCount, - ClientErrorCount: accumulator.clientErrorCount, - ServerErrorCount: accumulator.serverErrorCount, - }) - } - sortNodeAccessLogBucketRows(rows, query.SortBy, query.SortOrder) - return rows, nil + return items[start:end], nil } -func buildNodeAccessLogBucketIPRows(query NodeAccessLogBucketIPQuery) ([]*NodeAccessLogBucketIPRow, error) { - if query.BucketStartedAt.IsZero() { - return []*NodeAccessLogBucketIPRow{}, nil +func nodeAccessLogFetchLimit(page int, pageSize int) int { + if page < 0 { + page = 0 } - foldMinutes := query.FoldMinutes - if foldMinutes <= 0 { - foldMinutes = 3 + if pageSize <= 0 { + return 0 } - bucketStartedAt := query.BucketStartedAt.UTC() - logs, err := listNodeAccessLogsAcrossShards(NodeAccessLogQuery{ - NodeID: query.NodeID, - RemoteAddr: query.RemoteAddr, - Host: query.Host, - Path: query.Path, - Since: bucketStartedAt, - Until: bucketStartedAt.Add(time.Duration(foldMinutes) * time.Minute), - }) - if err != nil { - return nil, err - } - type accumulator struct { - requestCount int64 - successCount int64 - clientErrorCount int64 - serverErrorCount int64 - lastSeenAt time.Time - } - accumulators := make(map[string]*accumulator) - for _, item := range logs { - if item == nil { - continue - } - remoteAddr := strings.TrimSpace(item.RemoteAddr) - if remoteAddr == "" { - continue - } - acc := accumulators[remoteAddr] - if acc == nil { - acc = &accumulator{} - accumulators[remoteAddr] = acc - } - acc.requestCount++ - switch { - case item.StatusCode < 400: - acc.successCount++ - case item.StatusCode < 500: - acc.clientErrorCount++ - default: - acc.serverErrorCount++ - } - if item.LoggedAt.After(acc.lastSeenAt) { - acc.lastSeenAt = item.LoggedAt - } - } - rows := make([]*NodeAccessLogBucketIPRow, 0, len(accumulators)) - for remoteAddr, acc := range accumulators { - rows = append(rows, &NodeAccessLogBucketIPRow{ - RemoteAddr: remoteAddr, - RequestCount: acc.requestCount, - SuccessCount: acc.successCount, - ClientErrorCount: acc.clientErrorCount, - ServerErrorCount: acc.serverErrorCount, - LastSeenEpoch: acc.lastSeenAt.Unix(), - }) - } - sortNodeAccessLogBucketIPRows(rows, query.SortBy, query.SortOrder) - return rows, nil + return (page + 1) * pageSize } -func buildNodeAccessLogIPSummaryRows(query NodeAccessLogIPSummaryQuery, recentSince time.Time) ([]*NodeAccessLogIPSummaryRow, error) { - logs, err := listNodeAccessLogsAcrossShards(NodeAccessLogQuery{ - NodeID: query.NodeID, - RemoteAddr: query.RemoteAddr, - Host: query.Host, - Since: query.Since, - }) - if err != nil { - return nil, err +func nodeAccessLogOrderClause(sortBy string, sortOrder string) string { + direction := "DESC" + if normalizeSortOrder(sortOrder) == "asc" { + direction = "ASC" } - type accumulator struct { - totalRequests int64 - recentRequests int64 - lastSeenAt time.Time + column := "logged_at" + switch strings.TrimSpace(sortBy) { + case "status_code": + column = "status_code" + case "remote_addr": + column = "remote_addr" + case "host": + column = "host" + case "path": + column = "path" } - accumulators := make(map[string]*accumulator) - for _, item := range logs { - if item == nil { - continue - } - remoteAddr := strings.TrimSpace(item.RemoteAddr) - if remoteAddr == "" { - continue - } - acc := accumulators[remoteAddr] - if acc == nil { - acc = &accumulator{} - accumulators[remoteAddr] = acc - } - acc.totalRequests++ - if !recentSince.IsZero() && !item.LoggedAt.Before(recentSince) { - acc.recentRequests++ - } - if item.LoggedAt.After(acc.lastSeenAt) { - acc.lastSeenAt = item.LoggedAt - } + if column == "logged_at" { + return column + " " + direction + ", id " + direction } - rows := make([]*NodeAccessLogIPSummaryRow, 0, len(accumulators)) - for remoteAddr, acc := range accumulators { - rows = append(rows, &NodeAccessLogIPSummaryRow{ - RemoteAddr: remoteAddr, - TotalRequests: acc.totalRequests, - RecentRequests: acc.recentRequests, - LastSeenEpoch: acc.lastSeenAt.Unix(), - }) - } - sortNodeAccessLogIPSummaryRows(rows, query.SortBy, query.SortOrder) - return rows, nil + return column + " " + direction + ", logged_at " + direction + ", id " + direction } func sortNodeAccessLogBucketIPRows(items []*NodeAccessLogBucketIPRow, sortBy string, sortOrder string) { diff --git a/openflare-server/internal/model/node_access_log_agg.go b/openflare-server/internal/model/node_access_log_agg.go new file mode 100644 index 00000000..3cc0a074 --- /dev/null +++ b/openflare-server/internal/model/node_access_log_agg.go @@ -0,0 +1,421 @@ +package model + +import ( + "fmt" + "sort" + "strings" + "time" + + "gorm.io/gorm" +) + +type shardBucketAggregateRow struct { + BucketEpoch int64 `gorm:"column:bucket_epoch"` + RequestCount int64 `gorm:"column:request_count"` + SuccessCount int64 `gorm:"column:success_count"` + ClientErrorCount int64 `gorm:"column:client_error_count"` + ServerErrorCount int64 `gorm:"column:server_error_count"` +} + +type shardBucketDimensionRow struct { + BucketEpoch int64 `gorm:"column:bucket_epoch"` + Value string `gorm:"column:value"` +} + +type shardIPAggregateRow struct { + RemoteAddr string `gorm:"column:remote_addr"` + RequestCount int64 `gorm:"column:request_count"` + SuccessCount int64 `gorm:"column:success_count"` + ClientErrorCount int64 `gorm:"column:client_error_count"` + ServerErrorCount int64 `gorm:"column:server_error_count"` + LastSeenEpoch int64 `gorm:"column:last_seen_epoch"` +} + +type shardIPSummaryRow struct { + RemoteAddr string `gorm:"column:remote_addr"` + TotalRequests int64 `gorm:"column:total_requests"` + RecentRequests int64 `gorm:"column:recent_requests"` + LastSeenEpoch int64 `gorm:"column:last_seen_epoch"` +} + +type shardIPTrendRow struct { + BucketEpoch int64 `gorm:"column:bucket_epoch"` + RequestCount int64 `gorm:"column:request_count"` +} + +func buildNodeAccessLogBucketRows(query NodeAccessLogBucketQuery) ([]*NodeAccessLogBucketRow, error) { + db := normalizeShardedDB(DB) + filter := nodeAccessLogQueryFromBucket(query) + clause, args := buildNodeAccessLogFilterClause(filter) + bucketSeconds := int64(query.FoldMinutes * 60) + if bucketSeconds <= 0 { + bucketSeconds = 180 + } + bucketExpr := accessLogBucketEpochExpr(databaseDialect(db), bucketSeconds) + + type bucketAccumulator struct { + requestCount int64 + uniqueIPs map[string]struct{} + uniqueHosts map[string]struct{} + successCount int64 + clientErrorCount int64 + serverErrorCount int64 + } + accumulators := make(map[int64]*bucketAccumulator) + + for _, table := range observabilityShardTables("node_access_logs") { + var partials []shardBucketAggregateRow + sql := fmt.Sprintf(` +SELECT + %s AS bucket_epoch, + COUNT(*) AS request_count, + SUM(CASE WHEN status_code < 400 THEN 1 ELSE 0 END) AS success_count, + SUM(CASE WHEN status_code >= 400 AND status_code < 500 THEN 1 ELSE 0 END) AS client_error_count, + SUM(CASE WHEN status_code >= 500 THEN 1 ELSE 0 END) AS server_error_count +FROM %s +WHERE %s +GROUP BY bucket_epoch`, bucketExpr, table, clause) + if err := db.Raw(sql, args...).Scan(&partials).Error; err != nil { + return nil, err + } + for _, partial := range partials { + accumulator := accumulators[partial.BucketEpoch] + if accumulator == nil { + accumulator = &bucketAccumulator{ + uniqueIPs: make(map[string]struct{}), + uniqueHosts: make(map[string]struct{}), + } + accumulators[partial.BucketEpoch] = accumulator + } + accumulator.requestCount += partial.RequestCount + accumulator.successCount += partial.SuccessCount + accumulator.clientErrorCount += partial.ClientErrorCount + accumulator.serverErrorCount += partial.ServerErrorCount + } + + for _, column := range []string{"remote_addr", "host"} { + dimensions, err := queryBucketDimensionRows(db, table, clause, args, column, bucketExpr) + if err != nil { + return nil, err + } + for _, item := range dimensions { + accumulator := accumulators[item.BucketEpoch] + if accumulator == nil { + accumulator = &bucketAccumulator{ + uniqueIPs: make(map[string]struct{}), + uniqueHosts: make(map[string]struct{}), + } + accumulators[item.BucketEpoch] = accumulator + } + trimmed := strings.TrimSpace(item.Value) + if trimmed == "" { + continue + } + switch column { + case "remote_addr": + accumulator.uniqueIPs[trimmed] = struct{}{} + case "host": + accumulator.uniqueHosts[trimmed] = struct{}{} + } + } + } + } + + rows := make([]*NodeAccessLogBucketRow, 0, len(accumulators)) + for bucketEpoch, accumulator := range accumulators { + rows = append(rows, &NodeAccessLogBucketRow{ + BucketEpoch: bucketEpoch, + RequestCount: accumulator.requestCount, + UniqueIPCount: int64(len(accumulator.uniqueIPs)), + UniqueHostCount: int64(len(accumulator.uniqueHosts)), + SuccessCount: accumulator.successCount, + ClientErrorCount: accumulator.clientErrorCount, + ServerErrorCount: accumulator.serverErrorCount, + }) + } + sortNodeAccessLogBucketRows(rows, query.SortBy, query.SortOrder) + return rows, nil +} + +func queryBucketDimensionRows(db *gorm.DB, table string, clause string, args []any, column string, bucketExpr string) ([]shardBucketDimensionRow, error) { + var rows []shardBucketDimensionRow + sql := fmt.Sprintf(` +SELECT + %s AS bucket_epoch, + TRIM(%s) AS value +FROM %s +WHERE %s AND TRIM(%s) <> '' +GROUP BY bucket_epoch, TRIM(%s)`, bucketExpr, column, table, clause, column, column) + if err := db.Raw(sql, args...).Scan(&rows).Error; err != nil { + return nil, err + } + return rows, nil +} + +func buildNodeAccessLogBucketIPRows(query NodeAccessLogBucketIPQuery) ([]*NodeAccessLogBucketIPRow, error) { + if query.BucketStartedAt.IsZero() { + return []*NodeAccessLogBucketIPRow{}, nil + } + foldMinutes := query.FoldMinutes + if foldMinutes <= 0 { + foldMinutes = 3 + } + bucketStartedAt := query.BucketStartedAt.UTC() + filter := NodeAccessLogQuery{ + NodeID: query.NodeID, + RemoteAddr: query.RemoteAddr, + Host: query.Host, + Path: query.Path, + Since: bucketStartedAt, + Until: bucketStartedAt.Add(time.Duration(foldMinutes) * time.Minute), + } + rows, err := queryIPAggregateRows(filter, false) + if err != nil { + return nil, err + } + sortNodeAccessLogBucketIPRows(rows, query.SortBy, query.SortOrder) + return rows, nil +} + +func buildNodeAccessLogIPSummaryRows(query NodeAccessLogIPSummaryQuery, recentSince time.Time) ([]*NodeAccessLogIPSummaryRow, error) { + filter := NodeAccessLogQuery{ + NodeID: query.NodeID, + RemoteAddr: query.RemoteAddr, + Host: query.Host, + Since: query.Since, + } + db := normalizeShardedDB(DB) + clause, args := buildNodeAccessLogFilterClause(filter) + lastSeenExpr := accessLogEpochExpr(databaseDialect(db)) + + type accumulator struct { + totalRequests int64 + recentRequests int64 + lastSeenEpoch int64 + } + accumulators := make(map[string]*accumulator) + + for _, table := range observabilityShardTables("node_access_logs") { + recentClause := "0" + queryArgs := make([]any, 0, len(args)+1) + if !recentSince.IsZero() { + recentClause = "CASE WHEN logged_at >= ? THEN 1 ELSE 0 END" + queryArgs = append(queryArgs, recentSince) + } + queryArgs = append(queryArgs, args...) + var partials []shardIPSummaryRow + sql := fmt.Sprintf(` +SELECT + TRIM(remote_addr) AS remote_addr, + COUNT(*) AS total_requests, + SUM(%s) AS recent_requests, + MAX(%s) AS last_seen_epoch +FROM %s +WHERE %s AND TRIM(remote_addr) <> '' +GROUP BY TRIM(remote_addr)`, recentClause, lastSeenExpr, table, clause) + if err := db.Raw(sql, queryArgs...).Scan(&partials).Error; err != nil { + return nil, err + } + for _, partial := range partials { + remoteAddr := strings.TrimSpace(partial.RemoteAddr) + if remoteAddr == "" { + continue + } + acc := accumulators[remoteAddr] + if acc == nil { + acc = &accumulator{} + accumulators[remoteAddr] = acc + } + acc.totalRequests += partial.TotalRequests + acc.recentRequests += partial.RecentRequests + if partial.LastSeenEpoch > acc.lastSeenEpoch { + acc.lastSeenEpoch = partial.LastSeenEpoch + } + } + } + + rows := make([]*NodeAccessLogIPSummaryRow, 0, len(accumulators)) + for remoteAddr, acc := range accumulators { + rows = append(rows, &NodeAccessLogIPSummaryRow{ + RemoteAddr: remoteAddr, + TotalRequests: acc.totalRequests, + RecentRequests: acc.recentRequests, + LastSeenEpoch: acc.lastSeenEpoch, + }) + } + sortNodeAccessLogIPSummaryRows(rows, query.SortBy, query.SortOrder) + return rows, nil +} + +func queryIPAggregateRows(filter NodeAccessLogQuery, exactRemoteAddr bool) ([]*NodeAccessLogBucketIPRow, error) { + db := normalizeShardedDB(DB) + clause, args := buildNodeAccessLogFilterClause(filter) + lastSeenExpr := accessLogEpochExpr(databaseDialect(db)) + + type accumulator struct { + requestCount int64 + successCount int64 + clientErrorCount int64 + serverErrorCount int64 + lastSeenEpoch int64 + } + accumulators := make(map[string]*accumulator) + + for _, table := range observabilityShardTables("node_access_logs") { + queryClause := clause + queryArgs := append([]any{}, args...) + if exactRemoteAddr { + trimmed := strings.TrimSpace(filter.RemoteAddr) + if trimmed == "" { + return []*NodeAccessLogBucketIPRow{}, nil + } + queryClause = combineSQLClauses(queryClause, "TRIM(remote_addr) = ?") + queryArgs = append(queryArgs, trimmed) + } + var partials []shardIPAggregateRow + sql := fmt.Sprintf(` +SELECT + TRIM(remote_addr) AS remote_addr, + COUNT(*) AS request_count, + SUM(CASE WHEN status_code < 400 THEN 1 ELSE 0 END) AS success_count, + SUM(CASE WHEN status_code >= 400 AND status_code < 500 THEN 1 ELSE 0 END) AS client_error_count, + SUM(CASE WHEN status_code >= 500 THEN 1 ELSE 0 END) AS server_error_count, + MAX(%s) AS last_seen_epoch +FROM %s +WHERE %s AND TRIM(remote_addr) <> '' +GROUP BY TRIM(remote_addr)`, lastSeenExpr, table, queryClause) + if err := db.Raw(sql, queryArgs...).Scan(&partials).Error; err != nil { + return nil, err + } + for _, partial := range partials { + remoteAddr := strings.TrimSpace(partial.RemoteAddr) + if remoteAddr == "" { + continue + } + acc := accumulators[remoteAddr] + if acc == nil { + acc = &accumulator{} + accumulators[remoteAddr] = acc + } + acc.requestCount += partial.RequestCount + acc.successCount += partial.SuccessCount + acc.clientErrorCount += partial.ClientErrorCount + acc.serverErrorCount += partial.ServerErrorCount + if partial.LastSeenEpoch > acc.lastSeenEpoch { + acc.lastSeenEpoch = partial.LastSeenEpoch + } + } + } + + rows := make([]*NodeAccessLogBucketIPRow, 0, len(accumulators)) + for remoteAddr, acc := range accumulators { + rows = append(rows, &NodeAccessLogBucketIPRow{ + RemoteAddr: remoteAddr, + RequestCount: acc.requestCount, + SuccessCount: acc.successCount, + ClientErrorCount: acc.clientErrorCount, + ServerErrorCount: acc.serverErrorCount, + LastSeenEpoch: acc.lastSeenEpoch, + }) + } + return rows, nil +} + +func queryIPTrendRows(query NodeAccessLogIPTrendQuery) ([]*NodeAccessLogTrendPointRow, error) { + remoteAddr := strings.TrimSpace(query.RemoteAddr) + if remoteAddr == "" { + return []*NodeAccessLogTrendPointRow{}, nil + } + db := normalizeShardedDB(DB) + filter := NodeAccessLogQuery{ + NodeID: query.NodeID, + RemoteAddr: remoteAddr, + Host: query.Host, + Since: query.Since, + } + clause, args := buildNodeAccessLogFilterClause(filter) + bucketSeconds := int64(query.BucketMinutes * 60) + if bucketSeconds <= 0 { + bucketSeconds = 1800 + } + bucketExpr := accessLogBucketEpochExpr(databaseDialect(db), bucketSeconds) + queryClause := combineSQLClauses(clause, "TRIM(remote_addr) = ?") + queryArgs := append(append([]any{}, args...), remoteAddr) + + buckets := make(map[int64]int64) + for _, table := range observabilityShardTables("node_access_logs") { + var partials []shardIPTrendRow + sql := fmt.Sprintf(` +SELECT + %s AS bucket_epoch, + COUNT(*) AS request_count +FROM %s +WHERE %s +GROUP BY bucket_epoch`, bucketExpr, table, queryClause) + if err := db.Raw(sql, queryArgs...).Scan(&partials).Error; err != nil { + return nil, err + } + for _, partial := range partials { + buckets[partial.BucketEpoch] += partial.RequestCount + } + } + + items := make([]*NodeAccessLogTrendPointRow, 0, len(buckets)) + for bucketEpoch, requestCount := range buckets { + items = append(items, &NodeAccessLogTrendPointRow{ + BucketEpoch: bucketEpoch, + RequestCount: requestCount, + }) + } + sort.Slice(items, func(i int, j int) bool { + return items[i].BucketEpoch < items[j].BucketEpoch + }) + return items, nil +} + +func nodeAccessLogQueryFromBucket(query NodeAccessLogBucketQuery) NodeAccessLogQuery { + return NodeAccessLogQuery{ + NodeID: query.NodeID, + RemoteAddr: query.RemoteAddr, + Host: query.Host, + Path: query.Path, + Since: query.Since, + } +} + +func databaseDialect(db *gorm.DB) string { + if db == nil || db.Dialector == nil { + return "sqlite" + } + switch db.Dialector.Name() { + case "postgres": + return "postgres" + default: + return "sqlite" + } +} + +func accessLogBucketEpochExpr(dialect string, bucketSeconds int64) string { + switch dialect { + case "postgres": + return fmt.Sprintf("FLOOR(EXTRACT(EPOCH FROM logged_at AT TIME ZONE 'UTC') / %d) * %d", bucketSeconds, bucketSeconds) + default: + return fmt.Sprintf("(CAST(strftime('%%s', logged_at) AS INTEGER) / %d) * %d", bucketSeconds, bucketSeconds) + } +} + +func accessLogEpochExpr(dialect string) string { + switch dialect { + case "postgres": + return "FLOOR(EXTRACT(EPOCH FROM logged_at AT TIME ZONE 'UTC'))::bigint" + default: + return "CAST((julianday(logged_at) - 2440587.5) * 86400 AS INTEGER)" + } +} + +func combineSQLClauses(left string, right string) string { + if strings.TrimSpace(left) == "" || left == "TRUE" { + return right + } + return left + " AND " + right +} diff --git a/openflare-server/internal/model/node_access_log_test.go b/openflare-server/internal/model/node_access_log_test.go new file mode 100644 index 00000000..619cbd68 --- /dev/null +++ b/openflare-server/internal/model/node_access_log_test.go @@ -0,0 +1,542 @@ +package model + +import ( + "fmt" + "sort" + "strings" + "testing" + "time" +) + +func TestListNodeAccessLogsPaginatedAcrossShards(t *testing.T) { + db := openBareTestSQLiteDB(t, "node_access_log_pagination.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) + } + previousDB := DB + DB = db + t.Cleanup(func() { + DB = previousDB + }) + + now := time.Now().UTC() + for index := range 15 { + record := &NodeAccessLog{ + NodeID: "node-page", + LoggedAt: now.Add(-time.Duration(index) * time.Minute), + RemoteAddr: fmt.Sprintf("203.0.113.%d", (index%5)+1), + Host: "example.com", + Path: fmt.Sprintf("/path-%02d", index), + StatusCode: 200, + } + if err := db.Create(record).Error; err != nil { + t.Fatalf("seed access log %d: %v", index, err) + } + } + + query := NodeAccessLogQuery{ + NodeID: "node-page", + Page: 1, + PageSize: 5, + SortBy: "logged_at", + SortOrder: "desc", + } + page, err := ListNodeAccessLogs(query) + if err != nil { + t.Fatalf("ListNodeAccessLogs failed: %v", err) + } + if len(page) != 5 { + t.Fatalf("expected 5 rows, got %d", len(page)) + } + if page[0].Path != "/path-05" || page[4].Path != "/path-09" { + t.Fatalf("unexpected page ordering: %+v", page) + } + + totalRecords, totalIPs, err := CountNodeAccessLogs(query) + if err != nil { + t.Fatalf("CountNodeAccessLogs failed: %v", err) + } + if totalRecords != 15 { + t.Fatalf("expected total_records=15, got %d", totalRecords) + } + if totalIPs != 5 { + t.Fatalf("expected total_ip=5, got %d", totalIPs) + } +} + +func TestNodeAccessLogOptimizedQueriesMatchReference(t *testing.T) { + db := openBareTestSQLiteDB(t, "node_access_log_correctness.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) + } + previousDB := DB + DB = db + t.Cleanup(func() { + DB = previousDB + }) + + now := time.Now().UTC() + records := []*NodeAccessLog{ + {NodeID: "node-a", LoggedAt: now.Add(-5 * time.Minute), RemoteAddr: "1.1.1.1", Host: "a.example.com", Path: "/alpha", StatusCode: 200}, + {NodeID: "node-a", LoggedAt: now.Add(-4 * time.Minute), RemoteAddr: "2.2.2.2", Host: "a.example.com", Path: "/beta", StatusCode: 404}, + {NodeID: "node-b", LoggedAt: now.Add(-3 * time.Minute), RemoteAddr: "1.1.1.1", Host: "b.example.com", Path: "/gamma", StatusCode: 502}, + {NodeID: "node-b", LoggedAt: now.Add(-2 * time.Minute), RemoteAddr: " 3.3.3.3 ", Host: "b.example.com", Path: "/delta", StatusCode: 200}, + {NodeID: "node-b", LoggedAt: now.Add(-1 * time.Minute), RemoteAddr: "", Host: "b.example.com", Path: "/empty-ip", StatusCode: 200}, + } + for _, record := range records { + if err := db.Create(record).Error; err != nil { + t.Fatalf("seed access log: %v", err) + } + } + + baseQuery := NodeAccessLogQuery{ + Since: now.Add(-10 * time.Minute), + SortBy: "logged_at", + SortOrder: "desc", + } + reference, err := listNodeAccessLogsAcrossShards(baseQuery) + if err != nil { + t.Fatalf("reference list failed: %v", err) + } + referenceTotal, referenceIPs, err := countNodeAccessLogsReference(baseQuery) + if err != nil { + t.Fatalf("reference count failed: %v", err) + } + + totalRecords, totalIPs, err := CountNodeAccessLogs(baseQuery) + if err != nil { + t.Fatalf("CountNodeAccessLogs failed: %v", err) + } + if totalRecords != referenceTotal { + t.Fatalf("total_records mismatch: got %d want %d", totalRecords, referenceTotal) + } + if totalIPs != referenceIPs { + t.Fatalf("total_ip mismatch: got %d want %d", totalIPs, referenceIPs) + } + if totalRecords != int64(len(reference)) { + t.Fatalf("total_records should equal reference rows: got %d want %d", totalRecords, len(reference)) + } + + for page := range 3 { + query := baseQuery + query.Page = page + query.PageSize = 2 + pageRows, err := ListNodeAccessLogs(query) + if err != nil { + t.Fatalf("ListNodeAccessLogs page %d failed: %v", page, err) + } + start, end := paginateBounds(len(reference), page, query.PageSize) + if start >= len(reference) { + if len(pageRows) != 0 { + t.Fatalf("page %d expected empty slice, got %d rows", page, len(pageRows)) + } + continue + } + want := reference[start:end] + if !nodeAccessLogsEqual(pageRows, want) { + t.Fatalf("page %d mismatch:\n got=%+v\nwant=%+v", page, pageRows, want) + } + } + + filteredQuery := NodeAccessLogQuery{ + NodeID: "node-a", + Since: baseQuery.Since, + SortBy: "status_code", + SortOrder: "asc", + Page: 0, + PageSize: 10, + } + filteredReference, err := listNodeAccessLogsAcrossShards(filteredQuery) + if err != nil { + t.Fatalf("filtered reference list failed: %v", err) + } + filteredRows, err := ListNodeAccessLogs(filteredQuery) + if err != nil { + t.Fatalf("filtered ListNodeAccessLogs failed: %v", err) + } + if !nodeAccessLogsEqual(filteredRows, filteredReference) { + t.Fatalf("filtered list mismatch:\n got=%+v\nwant=%+v", filteredRows, filteredReference) + } + filteredTotal, filteredIPs, err := CountNodeAccessLogs(filteredQuery) + if err != nil { + t.Fatalf("filtered CountNodeAccessLogs failed: %v", err) + } + wantFilteredTotal, wantFilteredIPs, err := countNodeAccessLogsReference(filteredQuery) + if err != nil { + t.Fatalf("filtered reference count failed: %v", err) + } + if filteredTotal != wantFilteredTotal || filteredIPs != wantFilteredIPs { + t.Fatalf("filtered count mismatch: got (%d,%d) want (%d,%d)", filteredTotal, filteredIPs, wantFilteredTotal, wantFilteredIPs) + } +} + +func countNodeAccessLogsReference(query NodeAccessLogQuery) (int64, int64, error) { + all, err := listNodeAccessLogsAcrossShards(query) + if err != nil { + return 0, 0, err + } + ips := make(map[string]struct{}) + for _, item := range all { + if item == nil { + continue + } + if trimmed := strings.TrimSpace(item.RemoteAddr); trimmed != "" { + ips[trimmed] = struct{}{} + } + } + return int64(len(all)), int64(len(ips)), nil +} + +func nodeAccessLogsEqual(left []*NodeAccessLog, right []*NodeAccessLog) bool { + if len(left) != len(right) { + return false + } + for index := range left { + if left[index] == nil || right[index] == nil { + if left[index] != right[index] { + return false + } + continue + } + if left[index].ID != right[index].ID || + left[index].NodeID != right[index].NodeID || + !left[index].LoggedAt.Equal(right[index].LoggedAt) || + left[index].RemoteAddr != right[index].RemoteAddr || + left[index].Host != right[index].Host || + left[index].Path != right[index].Path || + left[index].StatusCode != right[index].StatusCode { + return false + } + } + return true +} + +func TestNodeAccessLogAggregationsMatchReference(t *testing.T) { + db := openBareTestSQLiteDB(t, "node_access_log_agg.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) + } + previousDB := DB + DB = db + t.Cleanup(func() { + DB = previousDB + }) + + now := time.Date(2026, 3, 19, 8, 12, 30, 0, time.UTC) + records := []*NodeAccessLog{ + {NodeID: "node-folded", LoggedAt: now.Add(-4 * time.Minute), RemoteAddr: "203.0.113.1", Host: "alpha.example.com", Path: "/first", StatusCode: 200}, + {NodeID: "node-folded", LoggedAt: now.Add(-3 * time.Minute), RemoteAddr: "203.0.113.1", Host: "alpha.example.com", Path: "/second", StatusCode: 502}, + {NodeID: "node-folded", LoggedAt: now.Add(-2 * time.Minute), RemoteAddr: "203.0.113.2", Host: "beta.example.com", Path: "/third", StatusCode: 404}, + } + for _, record := range records { + if err := db.Create(record).Error; err != nil { + t.Fatalf("seed access log: %v", err) + } + } + + since := now.Add(-10 * time.Minute) + bucketRows, err := buildNodeAccessLogBucketRows(NodeAccessLogBucketQuery{ + NodeID: "node-folded", Since: since, FoldMinutes: 5, SortBy: "request_count", SortOrder: "desc", + }) + if err != nil { + t.Fatalf("buildNodeAccessLogBucketRows failed: %v", err) + } + referenceBuckets := referenceBucketRows(records, 5, "request_count", "desc") + if !bucketRowsEqual(bucketRows, referenceBuckets) { + t.Fatalf("bucket rows mismatch:\n got=%+v\nwant=%+v", bucketRows, referenceBuckets) + } + + if len(bucketRows) == 0 { + t.Fatal("expected bucket rows before bucket ip verification") + } + bucketStartedAt := time.Unix(bucketRows[0].BucketEpoch, 0).UTC() + bucketIPRows, err := buildNodeAccessLogBucketIPRows(NodeAccessLogBucketIPQuery{ + NodeID: "node-folded", BucketStartedAt: bucketStartedAt, FoldMinutes: 5, SortBy: "request_count", SortOrder: "desc", + }) + if err != nil { + t.Fatalf("buildNodeAccessLogBucketIPRows failed: %v", err) + } + referenceBucketIPs := referenceBucketIPRows(records, bucketStartedAt, 5, "request_count", "desc") + if !bucketIPRowsEqual(bucketIPRows, referenceBucketIPs) { + if len(bucketIPRows) > 0 && len(referenceBucketIPs) > 0 { + t.Fatalf("bucket ip rows mismatch:\n got=%+v\nwant=%+v", *bucketIPRows[0], *referenceBucketIPs[0]) + } + t.Fatalf("bucket ip rows mismatch:\n got=%+v\nwant=%+v", bucketIPRows, referenceBucketIPs) + } + + recentSince := now.Add(-150 * time.Minute) + summaryRows, err := buildNodeAccessLogIPSummaryRows(NodeAccessLogIPSummaryQuery{ + NodeID: "node-folded", Since: since, SortBy: "total_requests", SortOrder: "desc", + }, recentSince) + if err != nil { + t.Fatalf("buildNodeAccessLogIPSummaryRows failed: %v", err) + } + referenceSummaries := referenceIPSummaryRows(records, since, recentSince, "total_requests", "desc") + if !ipSummaryRowsEqual(summaryRows, referenceSummaries) { + t.Fatalf("ip summary rows mismatch:\n got=%+v\nwant=%+v", summaryRows, referenceSummaries) + } + + trendRows, err := queryIPTrendRows(NodeAccessLogIPTrendQuery{ + NodeID: "node-folded", RemoteAddr: "203.0.113.1", Since: since, BucketMinutes: 5, + }) + if err != nil { + t.Fatalf("queryIPTrendRows failed: %v", err) + } + referenceTrend := referenceIPTrendRows(records, "203.0.113.1", 5) + if !trendRowsEqual(trendRows, referenceTrend) { + t.Fatalf("trend rows mismatch:\n got=%+v\nwant=%+v", trendRows, referenceTrend) + } +} + +func referenceBucketRows(records []*NodeAccessLog, foldMinutes int, sortBy string, sortOrder string) []*NodeAccessLogBucketRow { + type bucketAccumulator struct { + requestCount int64 + uniqueIPs map[string]struct{} + uniqueHosts map[string]struct{} + successCount int64 + clientErrorCount int64 + serverErrorCount int64 + } + accumulators := make(map[int64]*bucketAccumulator) + for _, item := range records { + if item == nil { + continue + } + bucketEpoch := bucketEpochForTime(item.LoggedAt, foldMinutes) + accumulator := accumulators[bucketEpoch] + if accumulator == nil { + accumulator = &bucketAccumulator{ + uniqueIPs: make(map[string]struct{}), + uniqueHosts: make(map[string]struct{}), + } + accumulators[bucketEpoch] = accumulator + } + accumulator.requestCount++ + if trimmed := strings.TrimSpace(item.RemoteAddr); trimmed != "" { + accumulator.uniqueIPs[trimmed] = struct{}{} + } + if trimmed := strings.TrimSpace(item.Host); trimmed != "" { + accumulator.uniqueHosts[trimmed] = struct{}{} + } + switch { + case item.StatusCode < 400: + accumulator.successCount++ + case item.StatusCode < 500: + accumulator.clientErrorCount++ + default: + accumulator.serverErrorCount++ + } + } + rows := make([]*NodeAccessLogBucketRow, 0, len(accumulators)) + for bucketEpoch, accumulator := range accumulators { + rows = append(rows, &NodeAccessLogBucketRow{ + BucketEpoch: bucketEpoch, + RequestCount: accumulator.requestCount, + UniqueIPCount: int64(len(accumulator.uniqueIPs)), + UniqueHostCount: int64(len(accumulator.uniqueHosts)), + SuccessCount: accumulator.successCount, + ClientErrorCount: accumulator.clientErrorCount, + ServerErrorCount: accumulator.serverErrorCount, + }) + } + sortNodeAccessLogBucketRows(rows, sortBy, sortOrder) + return rows +} + +func referenceBucketIPRows(records []*NodeAccessLog, bucketStartedAt time.Time, foldMinutes int, sortBy string, sortOrder string) []*NodeAccessLogBucketIPRow { + type accumulator struct { + requestCount int64 + successCount int64 + clientErrorCount int64 + serverErrorCount int64 + lastSeenAt time.Time + } + accumulators := make(map[string]*accumulator) + until := bucketStartedAt.Add(time.Duration(foldMinutes) * time.Minute) + for _, item := range records { + if item == nil || item.LoggedAt.Before(bucketStartedAt) || !item.LoggedAt.Before(until) { + continue + } + remoteAddr := strings.TrimSpace(item.RemoteAddr) + if remoteAddr == "" { + continue + } + acc := accumulators[remoteAddr] + if acc == nil { + acc = &accumulator{} + accumulators[remoteAddr] = acc + } + acc.requestCount++ + switch { + case item.StatusCode < 400: + acc.successCount++ + case item.StatusCode < 500: + acc.clientErrorCount++ + default: + acc.serverErrorCount++ + } + if item.LoggedAt.After(acc.lastSeenAt) { + acc.lastSeenAt = item.LoggedAt + } + } + rows := make([]*NodeAccessLogBucketIPRow, 0, len(accumulators)) + for remoteAddr, acc := range accumulators { + rows = append(rows, &NodeAccessLogBucketIPRow{ + RemoteAddr: remoteAddr, + RequestCount: acc.requestCount, + SuccessCount: acc.successCount, + ClientErrorCount: acc.clientErrorCount, + ServerErrorCount: acc.serverErrorCount, + LastSeenEpoch: acc.lastSeenAt.Unix(), + }) + } + sortNodeAccessLogBucketIPRows(rows, sortBy, sortOrder) + return rows +} + +func referenceIPSummaryRows(records []*NodeAccessLog, since time.Time, recentSince time.Time, sortBy string, sortOrder string) []*NodeAccessLogIPSummaryRow { + type accumulator struct { + totalRequests int64 + recentRequests int64 + lastSeenAt time.Time + } + accumulators := make(map[string]*accumulator) + for _, item := range records { + if item == nil || item.LoggedAt.Before(since) { + continue + } + remoteAddr := strings.TrimSpace(item.RemoteAddr) + if remoteAddr == "" { + continue + } + acc := accumulators[remoteAddr] + if acc == nil { + acc = &accumulator{} + accumulators[remoteAddr] = acc + } + acc.totalRequests++ + if !recentSince.IsZero() && !item.LoggedAt.Before(recentSince) { + acc.recentRequests++ + } + if item.LoggedAt.After(acc.lastSeenAt) { + acc.lastSeenAt = item.LoggedAt + } + } + rows := make([]*NodeAccessLogIPSummaryRow, 0, len(accumulators)) + for remoteAddr, acc := range accumulators { + rows = append(rows, &NodeAccessLogIPSummaryRow{ + RemoteAddr: remoteAddr, + TotalRequests: acc.totalRequests, + RecentRequests: acc.recentRequests, + LastSeenEpoch: acc.lastSeenAt.Unix(), + }) + } + sortNodeAccessLogIPSummaryRows(rows, sortBy, sortOrder) + return rows +} + +func referenceIPTrendRows(records []*NodeAccessLog, remoteAddr string, bucketMinutes int) []*NodeAccessLogTrendPointRow { + buckets := make(map[int64]int64) + for _, item := range records { + if item == nil || strings.TrimSpace(item.RemoteAddr) != remoteAddr { + continue + } + buckets[bucketEpochForTime(item.LoggedAt, bucketMinutes)]++ + } + rows := make([]*NodeAccessLogTrendPointRow, 0, len(buckets)) + for bucketEpoch, requestCount := range buckets { + rows = append(rows, &NodeAccessLogTrendPointRow{BucketEpoch: bucketEpoch, RequestCount: requestCount}) + } + sort.Slice(rows, func(i int, j int) bool { return rows[i].BucketEpoch < rows[j].BucketEpoch }) + return rows +} + +func bucketRowsEqual(left []*NodeAccessLogBucketRow, right []*NodeAccessLogBucketRow) bool { + if len(left) != len(right) { + return false + } + for index := range left { + if left[index] == nil || right[index] == nil { + if left[index] != right[index] { + return false + } + continue + } + if *left[index] != *right[index] { + return false + } + } + return true +} + +func bucketIPRowsEqual(left []*NodeAccessLogBucketIPRow, right []*NodeAccessLogBucketIPRow) bool { + if len(left) != len(right) { + return false + } + for index := range left { + if left[index] == nil || right[index] == nil { + if left[index] != right[index] { + return false + } + continue + } + if *left[index] != *right[index] { + return false + } + } + return true +} + +func ipSummaryRowsEqual(left []*NodeAccessLogIPSummaryRow, right []*NodeAccessLogIPSummaryRow) bool { + if len(left) != len(right) { + return false + } + for index := range left { + if left[index] == nil || right[index] == nil { + if left[index] != right[index] { + return false + } + continue + } + if *left[index] != *right[index] { + return false + } + } + return true +} + +func trendRowsEqual(left []*NodeAccessLogTrendPointRow, right []*NodeAccessLogTrendPointRow) bool { + if len(left) != len(right) { + return false + } + for index := range left { + if left[index] == nil || right[index] == nil { + if left[index] != right[index] { + return false + } + continue + } + if *left[index] != *right[index] { + return false + } + } + return true +} + +func TestNodeAccessLogOrderClauseMatchesSort(t *testing.T) { + if got := nodeAccessLogOrderClause("logged_at", "desc"); got != "logged_at DESC, id DESC" { + t.Fatalf("unexpected logged_at order clause: %q", got) + } + if got := nodeAccessLogOrderClause("status_code", "asc"); got != "status_code ASC, logged_at ASC, id ASC" { + t.Fatalf("unexpected status_code order clause: %q", got) + } +}