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) } }