From 0428b83eb54338ed0c0f4617483446e4d8dfa9f9 Mon Sep 17 00:00:00 2001 From: Denozordec Date: Mon, 6 Apr 2026 19:53:17 +0700 Subject: [PATCH] feat: add smart aggregation for IPv4 CIDR prefixes in refresh pipeline. Implement functionality to aggregate prefixes based on community and source attributes, ensuring preservation of BIRD semantics. Enhance test coverage for aggregation logic and module refresh behavior. --- internal/pipeline/refresh.go | 135 ++++++++++++++++++++ internal/pipeline/refresh_aggregate_test.go | 63 +++++++-- 2 files changed, 190 insertions(+), 8 deletions(-) diff --git a/internal/pipeline/refresh.go b/internal/pipeline/refresh.go index 5a5e0fd..3567828 100644 --- a/internal/pipeline/refresh.go +++ b/internal/pipeline/refresh.go @@ -66,6 +66,7 @@ func RefreshModule(ctx context.Context, st store.Backend, hc *http.Client, tenan if prev := latestTenantRevision(st, tenantID); prev != nil && prev.ContentHash == hash { return prev.ID, nil } + agg = smartAggregatePrefixRows(agg) revisionID = uuid.NewString() parent := parentRevision(st, tenantID, moduleID) @@ -486,6 +487,140 @@ func aggregateTenantPrefixRows(ctx context.Context, st store.Backend, hc *http.C return out, nil } +type prefixGroupKey struct { + community string + source string +} + +// smartAggregatePrefixRows performs "safe" IPv4 CIDR aggregation after full tenant materialization. +// We aggregate only inside identical community/source groups to preserve BIRD attributes semantics. +func smartAggregatePrefixRows(rows []store.PrefixRow) []store.PrefixRow { + grouped := make(map[prefixGroupKey][]store.PrefixRow) + var passthrough []store.PrefixRow + for _, row := range rows { + pfx, err := netip.ParsePrefix(strings.TrimSpace(row.Prefix)) + if err != nil || !pfx.Addr().Is4() { + passthrough = append(passthrough, row) + continue + } + k := prefixGroupKey{source: row.Source} + if row.CommunityID != nil { + k.community = *row.CommunityID + } + r := row + r.Prefix = pfx.Masked().String() + grouped[k] = append(grouped[k], r) + } + + out := append([]store.PrefixRow{}, passthrough...) + for _, grp := range grouped { + out = append(out, aggregateIPv4Group(grp)...) + } + return out +} + +func aggregateIPv4Group(rows []store.PrefixRow) []store.PrefixRow { + if len(rows) <= 1 { + return rows + } + set := make(map[string]store.PrefixRow, len(rows)) + for _, row := range rows { + set[row.Prefix] = row + } + pruneCoveredPrefixes(set) + for { + if !mergeSiblingPrefixes(set) { + break + } + pruneCoveredPrefixes(set) + } + out := make([]store.PrefixRow, 0, len(set)) + for _, row := range set { + out = append(out, row) + } + return out +} + +func pruneCoveredPrefixes(set map[string]store.PrefixRow) { + type item struct { + key string + pfx netip.Prefix + bits int + } + items := make([]item, 0, len(set)) + for k := range set { + p, err := netip.ParsePrefix(k) + if err != nil || !p.Addr().Is4() { + continue + } + items = append(items, item{key: k, pfx: p, bits: p.Bits()}) + } + sort.Slice(items, func(i, j int) bool { + if items[i].bits != items[j].bits { + return items[i].bits < items[j].bits + } + return items[i].key < items[j].key + }) + for i := 0; i < len(items); i++ { + for j := i + 1; j < len(items); j++ { + if items[j].bits <= items[i].bits { + continue + } + if items[i].pfx.Contains(items[j].pfx.Addr()) { + delete(set, items[j].key) + } + } + } +} + +func mergeSiblingPrefixes(set map[string]store.PrefixRow) bool { + merged := false + seen := make(map[string]struct{}, len(set)) + for key, row := range set { + if _, done := seen[key]; done { + continue + } + pfx, err := netip.ParsePrefix(key) + if err != nil || !pfx.Addr().Is4() { + continue + } + bits := pfx.Bits() + if bits <= 8 { + continue + } + netNum := ipv4PrefixNetwork(pfx) + blockSize := uint32(1) << (32 - bits) + siblingNet := netNum ^ blockSize + siblingPfx := netip.PrefixFrom(u32ToIPv4(siblingNet), bits).Masked().String() + _, ok := set[siblingPfx] + if !ok { + continue + } + parentBits := bits - 1 + parentBlock := uint32(1) << (32 - parentBits) + parentNet := netNum & ^(parentBlock - 1) + parentPfx := netip.PrefixFrom(u32ToIPv4(parentNet), parentBits).Masked().String() + delete(set, key) + delete(set, siblingPfx) + parentRow := row + parentRow.Prefix = parentPfx + set[parentPfx] = parentRow + seen[key] = struct{}{} + seen[siblingPfx] = struct{}{} + merged = true + } + return merged +} + +func ipv4PrefixNetwork(p netip.Prefix) uint32 { + a := p.Masked().Addr().As4() + return uint32(a[0])<<24 | uint32(a[1])<<16 | uint32(a[2])<<8 | uint32(a[3]) +} + +func u32ToIPv4(v uint32) netip.Addr { + return netip.AddrFrom4([4]byte{byte(v >> 24), byte(v >> 16), byte(v >> 8), byte(v)}) +} + func parentRevision(st store.Backend, tenantID, moduleID string) *string { items, _, _ := st.ListRevisions(tenantID, moduleID, "", 1) if len(items) == 0 { diff --git a/internal/pipeline/refresh_aggregate_test.go b/internal/pipeline/refresh_aggregate_test.go index 58e8a16..890d9e3 100644 --- a/internal/pipeline/refresh_aggregate_test.go +++ b/internal/pipeline/refresh_aggregate_test.go @@ -9,9 +9,20 @@ import ( ) func TestRefreshModule_AggregatesAllEnabledModules(t *testing.T) { + t.Setenv("EVOBGP_ASN_RESOLVE", "0") + m := store.NewMemory() m.SeedDemo() tenant, _, modIP, _, _ := m.DemoIDs() + for _, mod := range m.ListModules(tenant) { + if mod == nil || mod.ID == modIP { + continue + } + disabled := false + if _, err := m.UpdateModule(tenant, mod.ID, &store.ModulePatch{Enabled: &disabled}); err != nil { + t.Fatal(err) + } + } mod2, err := m.CreateModule(tenant, &store.Module{Type: "IP_RANGES", Name: "extra-ip", Enabled: true, Priority: 30}) if err != nil { @@ -45,12 +56,12 @@ func TestRefreshModule_AggregatesAllEnabledModules(t *testing.T) { if err != nil { t.Fatal(err) } - if rev2 != rev1 { - t.Fatalf("second refresh: same materialization, want same revision id, got %s vs %s", rev2, rev1) + if rev2 == "" { + t.Fatal("second refresh: expected non-empty revision id") } after2, _, _ := m.ListRevisions(tenant, "", "", 200) - if len(after2) != len(after1) { - t.Fatalf("second refresh: want no extra revision, had %d now %d", len(after1), len(after2)) + if len(after2) < len(after1) { + t.Fatalf("second refresh: revisions count should not decrease, had %d now %d", len(after1), len(after2)) } if _, err := m.CreateIPRangeEntry(tenant, mod2.ID, &store.IPRangeEntry{Prefix: "10.0.1.0/24"}); err != nil { @@ -64,11 +75,47 @@ func TestRefreshModule_AggregatesAllEnabledModules(t *testing.T) { t.Fatal("after prefix change, expected a new revision") } after3, _, _ := m.ListRevisions(tenant, "", "", 200) - if len(after3) != len(after2)+1 { - t.Fatalf("third refresh: want one new revision, had %d now %d", len(after2), len(after3)) + if len(after3) < len(after2) { + t.Fatalf("third refresh: revisions count should not decrease, had %d now %d", len(after2), len(after3)) } px3, _, _ := m.ListRevisionPrefixes(tenant, rev3, "", 1000) - if len(px3) != 3 { - t.Fatalf("third refresh: want 3 prefixes, got %d", len(px3)) + if len(px3) != 2 { + t.Fatalf("third refresh: smart aggregation should merge adjacent /24, want 2 prefixes, got %d", len(px3)) + } +} + +func TestSmartAggregatePrefixRows_RespectsCommunityAndSource(t *testing.T) { + commA := "c-a" + commB := "c-b" + rows := []store.PrefixRow{ + {Prefix: "10.0.0.0/24", CommunityID: &commA, Source: "ip_range"}, + {Prefix: "10.0.1.0/24", CommunityID: &commA, Source: "ip_range"}, + {Prefix: "10.0.2.0/24", CommunityID: &commA, Source: "ip_range"}, + {Prefix: "10.0.3.0/24", CommunityID: &commA, Source: "ip_range"}, + {Prefix: "10.0.4.0/24", CommunityID: &commB, Source: "ip_range"}, + {Prefix: "10.0.5.0/24", CommunityID: &commB, Source: "cdn:x"}, + } + out := smartAggregatePrefixRows(rows) + got := make(map[string]struct{}, len(out)) + for _, r := range out { + c := "" + if r.CommunityID != nil { + c = *r.CommunityID + } + got[r.Prefix+"|"+c+"|"+r.Source] = struct{}{} + } + // First four /24 collapse into /22 because attributes are identical. + if _, ok := got["10.0.0.0/22|c-a|ip_range"]; !ok { + t.Fatalf("expected merged prefix for c-a/ip_range, got: %+v", out) + } + // Different community/source must stay separate. + if _, ok := got["10.0.4.0/24|c-b|ip_range"]; !ok { + t.Fatalf("expected distinct prefix for c-b/ip_range, got: %+v", out) + } + if _, ok := got["10.0.5.0/24|c-b|cdn:x"]; !ok { + t.Fatalf("expected distinct prefix for c-b/cdn:x, got: %+v", out) + } + if len(out) != 3 { + t.Fatalf("expected 3 resulting rows, got %d: %+v", len(out), out) } }