diff --git a/internal/pkg/table/policy.go b/internal/pkg/table/policy.go index 3377a6172..560bbd2ac 100644 --- a/internal/pkg/table/policy.go +++ b/internal/pkg/table/policy.go @@ -4646,27 +4646,23 @@ func (r *RoutingPolicy) Initialize() error { return nil } -func (r *RoutingPolicy) setPeerPolicy(id string, c oc.ApplyPolicy) { +func (r *RoutingPolicy) setPeerPolicy(id string, c oc.ApplyPolicy) error { for _, dir := range []PolicyDirection{POLICY_DIRECTION_IMPORT, POLICY_DIRECTION_EXPORT} { ps, def, err := r.getAssignmentFromConfig(dir, c) if err != nil { - r.logger.Error("failed to get policy info", - slog.String("Topic", "Policy"), - slog.String("Dir", dir.String()), - slog.String("Error", err.Error())) - continue + return fmt.Errorf("failed to get %s policy info for %s: %w", dir, id, err) } r.setDefaultPolicy(id, dir, def) r.setPolicy(id, dir, ps) } + return nil } func (r *RoutingPolicy) SetPeerPolicy(peerId string, c oc.ApplyPolicy) error { r.mu.Lock() defer r.mu.Unlock() - r.setPeerPolicy(peerId, c) - return nil + return r.setPeerPolicy(peerId, c) } func (r *RoutingPolicy) Reset(rp *oc.RoutingPolicy, ap map[string]oc.ApplyPolicy) error { @@ -4685,7 +4681,9 @@ func (r *RoutingPolicy) Reset(rp *oc.RoutingPolicy, ap map[string]oc.ApplyPolicy } for id, c := range ap { - r.setPeerPolicy(id, c) + if err := r.setPeerPolicy(id, c); err != nil { + return err + } } return nil } diff --git a/pkg/server/server.go b/pkg/server/server.go index d0b2b136a..f2c75115f 100644 --- a/pkg/server/server.go +++ b/pkg/server/server.go @@ -2261,6 +2261,41 @@ func (s *BgpServer) SetPolicies(ctx context.Context, r *api.SetPoliciesRequest) c.Config.DefaultExportPolicy = rt return c, nil } + applyAssignment := func(ap map[string]oc.ApplyPolicy, assignment *api.PolicyAssignment) error { + if assignment == nil { + return fmt.Errorf("nil policy assignment") + } + id, dir, err := s.toPolicyInfo(assignment.Name, assignment.Direction) + if err != nil { + return err + } + + policies := make([]string, 0, len(assignment.Policies)) + for _, policy := range assignment.Policies { + if policy == nil { + return fmt.Errorf("nil policy in assignment %q", assignment.Name) + } + policies = append(policies, policy.Name) + } + + c := ap[id] + switch dir { + case table.POLICY_DIRECTION_IMPORT: + c.Config.ImportPolicyList = policies + if defaultPolicy, ok := defaultPolicyTypeFromRouteAction(assignment.DefaultAction); ok { + c.Config.DefaultImportPolicy = defaultPolicy + } + case table.POLICY_DIRECTION_EXPORT: + c.Config.ExportPolicyList = policies + if defaultPolicy, ok := defaultPolicyTypeFromRouteAction(assignment.DefaultAction); ok { + c.Config.DefaultExportPolicy = defaultPolicy + } + default: + return fmt.Errorf("invalid policy direction") + } + ap[id] = c + return nil + } return s.mgmtOperation(func() error { ap := make(map[string]oc.ApplyPolicy, len(s.neighborMap)+1) @@ -2278,10 +2313,26 @@ func (s *BgpServer) SetPolicies(ctx context.Context, r *api.SetPoliciesRequest) } ap[peer.ID()] = *a } + for _, assignment := range r.Assignments { + if err := applyAssignment(ap, assignment); err != nil { + return err + } + } return s.policy.Reset(rp, ap) }, false) } +func defaultPolicyTypeFromRouteAction(action api.RouteAction) (oc.DefaultPolicyType, bool) { + switch action { + case api.RouteAction_ROUTE_ACTION_ACCEPT: + return oc.DEFAULT_POLICY_TYPE_ACCEPT_ROUTE, true + case api.RouteAction_ROUTE_ACTION_REJECT: + return oc.DEFAULT_POLICY_TYPE_REJECT_ROUTE, true + default: + return "", false + } +} + // EVPN MAC MOBILITY HANDLING // // We don't have multihoming function now, so ignore diff --git a/pkg/server/server_test.go b/pkg/server/server_test.go index 2b41ee6d0..773c9fdb9 100644 --- a/pkg/server/server_test.go +++ b/pkg/server/server_test.go @@ -496,6 +496,171 @@ func TestModPolicyAssign(t *testing.T) { assert.Equal(len(ps), 2) } +func TestSetPoliciesAssignments(t *testing.T) { + policy := &api.Policy{Name: "p1"} + validAssignments := []*api.PolicyAssignment{ + { + Name: table.GLOBAL_RIB_NAME, + Direction: api.PolicyDirection_POLICY_DIRECTION_IMPORT, + Policies: []*api.Policy{{Name: policy.Name}}, + DefaultAction: api.RouteAction_ROUTE_ACTION_REJECT, + }, + { + Name: table.GLOBAL_RIB_NAME, + Direction: api.PolicyDirection_POLICY_DIRECTION_EXPORT, + Policies: []*api.Policy{{Name: policy.Name}}, + DefaultAction: api.RouteAction_ROUTE_ACTION_REJECT, + }, + } + tests := []struct { + name string + request *api.SetPoliciesRequest + nextRequest *api.SetPoliciesRequest + err string + expectAssignments bool + }{ + { + name: "applies assignments", + request: &api.SetPoliciesRequest{ + Policies: []*api.Policy{policy}, + Assignments: validAssignments, + }, + expectAssignments: true, + }, + { + name: "missing policy", + request: &api.SetPoliciesRequest{ + Policies: []*api.Policy{policy}, + Assignments: []*api.PolicyAssignment{{ + Name: table.GLOBAL_RIB_NAME, + Direction: api.PolicyDirection_POLICY_DIRECTION_IMPORT, + Policies: []*api.Policy{{Name: "missing"}}, + DefaultAction: api.RouteAction_ROUTE_ACTION_REJECT, + }}, + }, + err: "not found policy missing", + }, + { + name: "duplicate policy", + request: &api.SetPoliciesRequest{ + Policies: []*api.Policy{policy}, + Assignments: []*api.PolicyAssignment{{ + Name: table.GLOBAL_RIB_NAME, + Direction: api.PolicyDirection_POLICY_DIRECTION_IMPORT, + Policies: []*api.Policy{{Name: policy.Name}, {Name: policy.Name}}, + DefaultAction: api.RouteAction_ROUTE_ACTION_REJECT, + }}, + }, + err: "duplicated policy p1", + }, + { + name: "nil assignment", + request: &api.SetPoliciesRequest{ + Policies: []*api.Policy{policy}, + Assignments: []*api.PolicyAssignment{nil}, + }, + err: "nil policy assignment", + }, + { + name: "nil policy", + request: &api.SetPoliciesRequest{ + Policies: []*api.Policy{policy}, + Assignments: []*api.PolicyAssignment{{ + Name: table.GLOBAL_RIB_NAME, + Direction: api.PolicyDirection_POLICY_DIRECTION_IMPORT, + Policies: []*api.Policy{nil}, + DefaultAction: api.RouteAction_ROUTE_ACTION_REJECT, + }}, + }, + err: "nil policy in assignment", + }, + { + name: "without assignments preserves existing assignments", + request: &api.SetPoliciesRequest{ + Policies: []*api.Policy{policy}, + Assignments: validAssignments, + }, + nextRequest: &api.SetPoliciesRequest{ + Policies: []*api.Policy{policy}, + }, + expectAssignments: true, + }, + { + name: "empty assignment policies clears existing policies", + request: &api.SetPoliciesRequest{ + Policies: []*api.Policy{policy}, + Assignments: validAssignments, + }, + nextRequest: &api.SetPoliciesRequest{ + Policies: []*api.Policy{policy}, + Assignments: []*api.PolicyAssignment{ + { + Name: table.GLOBAL_RIB_NAME, + Direction: api.PolicyDirection_POLICY_DIRECTION_IMPORT, + DefaultAction: api.RouteAction_ROUTE_ACTION_REJECT, + }, + { + Name: table.GLOBAL_RIB_NAME, + Direction: api.PolicyDirection_POLICY_DIRECTION_EXPORT, + DefaultAction: api.RouteAction_ROUTE_ACTION_REJECT, + }, + }, + }, + expectAssignments: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + require := require.New(t) + + s := NewBgpServer() + go s.Serve() + err := s.StartBgp(context.Background(), &api.StartBgpRequest{ + Global: &api.Global{ + Asn: 1, + RouterId: "1.1.1.1", + ListenPort: -1, + }, + }) + require.NoError(err) + defer s.StopBgp(context.Background(), &api.StopBgpRequest{}) + + err = s.SetPolicies(context.Background(), tt.request) + if tt.err != "" { + require.ErrorContains(err, tt.err) + return + } + require.NoError(err) + + if tt.nextRequest != nil { + err = s.SetPolicies(context.Background(), tt.nextRequest) + require.NoError(err) + } + + for _, direction := range []api.PolicyDirection{ + api.PolicyDirection_POLICY_DIRECTION_IMPORT, + api.PolicyDirection_POLICY_DIRECTION_EXPORT, + } { + assignments := []*api.PolicyAssignment{} + err = s.ListPolicyAssignment(context.Background(), &api.ListPolicyAssignmentRequest{ + Name: table.GLOBAL_RIB_NAME, + Direction: direction, + }, func(p *api.PolicyAssignment) { assignments = append(assignments, p) }) + require.NoError(err) + require.Len(assignments, 1) + require.Equal(api.RouteAction_ROUTE_ACTION_REJECT, assignments[0].DefaultAction) + if tt.expectAssignments { + require.Len(assignments[0].Policies, 1) + require.Equal(policy.Name, assignments[0].Policies[0].Name) + } else { + require.Empty(assignments[0].Policies) + } + } + }) + } +} + func TestBMPMonitoringPolicyFromAPI(t *testing.T) { t.Parallel()