Files
netbird-kubernetes-operator/internal/controller/nbpolicy_controller.go
T
Philip LaineandGitHub 7acd175882 Gateway API support (#117)
This change adds support for the new proxy service to the operator
through Gateway API. This change attempts to standardize concepts around
the Gateway API to allow for compatibility with other projects.

Fixes #111
Fixes #44

Signed-off-by: Philip Laine <philip.laine@gmail.com>
2026-03-19 13:01:58 +01:00

402 lines
14 KiB
Go

package controller
import (
"context"
"fmt"
"slices"
"strconv"
"strings"
"github.com/go-logr/logr"
netbird "github.com/netbirdio/netbird/shared/management/client/rest"
"github.com/netbirdio/netbird/shared/management/http/api"
"k8s.io/apimachinery/pkg/api/errors"
"k8s.io/apimachinery/pkg/types"
ctrl "sigs.k8s.io/controller-runtime"
"sigs.k8s.io/controller-runtime/pkg/client"
netbirdiov1 "github.com/netbirdio/kubernetes-operator/api/v1"
"github.com/netbirdio/kubernetes-operator/internal/util"
)
// NBPolicyReconciler reconciles a NBPolicy object
type NBPolicyReconciler struct {
client.Client
Netbird *netbird.Client
}
var (
errUnknownProtocol = fmt.Errorf("unknown protocol")
errKubernetesAPI = fmt.Errorf("kubernetes API error")
errNetBirdAPI = fmt.Errorf("netbird API error")
errInvalidValue = fmt.Errorf("invalid value")
)
const (
protocolTCP = "tcp"
protocolUDP = "udp"
)
// getResources get all NBResource objects in policy.status.managedServiceList
func (r *NBPolicyReconciler) getResources(ctx context.Context, nbPolicy *netbirdiov1.NBPolicy, logger logr.Logger) ([]netbirdiov1.NBResource, error) {
var resourceList []netbirdiov1.NBResource
var updatedManagedServiceList []string
for _, rss := range nbPolicy.Status.ManagedServiceList {
var resource netbirdiov1.NBResource
namespacedName := types.NamespacedName{Namespace: strings.Split(rss, "/")[0], Name: strings.Split(rss, "/")[1]}
err := r.Client.Get(ctx, namespacedName, &resource)
if err != nil && !errors.IsNotFound(err) {
logger.Error(errKubernetesAPI, "Error getting NBResource", "namespace", namespacedName.Namespace, "name", namespacedName.Name)
nbPolicy.Status.Conditions = netbirdiov1.NBConditionFalse("internalError", fmt.Sprintf("Error getting NBResource: %v", err))
return nil, err
}
if err == nil && resource.DeletionTimestamp == nil {
updatedManagedServiceList = append(updatedManagedServiceList, rss)
resourceList = append(resourceList, resource)
}
}
nbPolicy.Status.ManagedServiceList = updatedManagedServiceList
return resourceList, nil
}
// mapResources map each NBResource ports and protocols into one object to generate the policy
// returns map[protocol] => ports, destination group IDs
func (r *NBPolicyReconciler) mapResources(ctx context.Context, nbPolicy *netbirdiov1.NBPolicy, resources []netbirdiov1.NBResource, logger logr.Logger) (map[string][]int32, []string, error) {
portMapping := map[string]map[int32]any{
protocolTCP: make(map[int32]any),
protocolUDP: make(map[int32]any),
}
groups, err := r.groupNamesToIDs(ctx, nbPolicy.Spec.DestinationGroups, logger)
if err != nil {
return nil, nil, err
}
for _, resource := range resources {
generatedBy := nbPolicy.Annotations["netbird.io/generated-by"]
generatedBy = strings.ReplaceAll(generatedBy, "/", "-")
if resource.Status.PolicyName == nil {
continue
}
resourcePolicies := util.SplitTrim(*resource.Status.PolicyName, ",")
if generatedBy == "" && !slices.Contains(resourcePolicies, nbPolicy.Name) {
continue
}
if generatedBy != "" && !slices.Contains(resourcePolicies, strings.ReplaceAll(nbPolicy.Name, "-"+generatedBy, "")) {
continue
}
// Groups
groups = append(groups, resource.Status.Groups...)
for _, p := range resource.Spec.TCPPorts {
portMapping[protocolTCP][p] = nil
}
for _, p := range resource.Spec.UDPPorts {
portMapping[protocolUDP][p] = nil
}
}
ports := make(map[string][]int32)
for k, vs := range portMapping {
ports[k] = nil
for v := range vs {
ports[k] = append(ports[k], v)
}
slices.Sort(ports[k])
}
return ports, groups, nil
}
// createPolicy helper for creating policy with settings
func (r *NBPolicyReconciler) createPolicy(ctx context.Context, nbPolicy *netbirdiov1.NBPolicy, protocol string, sourceGroupIDs, destinationGroupIDs, ports []string, logger logr.Logger) (*string, error) {
policyName := fmt.Sprintf("%s %s", nbPolicy.Spec.Name, strings.ToUpper(protocol))
logger.Info("Creating NetBird Policy", "name", policyName, "description", nbPolicy.Spec.Description, "protocol", protocol, "sources", sourceGroupIDs, "destinations", destinationGroupIDs, "ports", ports, "bidirectional", nbPolicy.Spec.Bidirectional)
policy, err := r.Netbird.Policies.Create(ctx, api.PostApiPoliciesJSONRequestBody{
Enabled: true,
Name: policyName,
Description: &nbPolicy.Spec.Description,
Rules: []api.PolicyRuleUpdate{
{
Enabled: true,
Name: policyName,
Description: &nbPolicy.Spec.Description,
Action: api.PolicyRuleUpdateActionAccept,
Protocol: api.PolicyRuleUpdateProtocol(protocol),
Bidirectional: nbPolicy.Spec.Bidirectional,
Sources: &sourceGroupIDs,
Destinations: &destinationGroupIDs,
Ports: &ports,
},
},
})
if err != nil {
logger.Error(errNetBirdAPI, "Error creating Policy", "err", err)
nbPolicy.Status.Conditions = netbirdiov1.NBConditionFalse("APIError", fmt.Sprintf("Error creating policy: %v", err))
return nil, err
}
return policy.Id, nil
}
// updatePolicy helper for updating policy with settings
func (r *NBPolicyReconciler) updatePolicy(ctx context.Context, policyID *string, nbPolicy *netbirdiov1.NBPolicy, protocol string, sourceGroupIDs, destinationGroupIDs, ports []string, logger logr.Logger) (*string, bool, error) {
policyName := fmt.Sprintf("%s %s", nbPolicy.Spec.Name, strings.ToUpper(protocol))
logger.Info("Updating NetBird Policy", "name", policyName, "description", nbPolicy.Spec.Description, "protocol", protocol, "sources", sourceGroupIDs, "destinations", destinationGroupIDs, "ports", ports, "bidirectional", nbPolicy.Spec.Bidirectional)
_, err := r.Netbird.Policies.Update(ctx, *policyID, api.PutApiPoliciesPolicyIdJSONRequestBody{
Enabled: true,
Name: policyName,
Description: &nbPolicy.Spec.Description,
Rules: []api.PolicyRuleUpdate{
{
Enabled: true,
Name: policyName,
Description: &nbPolicy.Spec.Description,
Action: api.PolicyRuleUpdateActionAccept,
Protocol: api.PolicyRuleUpdateProtocol(protocol),
Bidirectional: nbPolicy.Spec.Bidirectional,
Sources: &sourceGroupIDs,
Destinations: &destinationGroupIDs,
Ports: &ports,
},
},
})
if err != nil && !strings.Contains(err.Error(), "not found") {
logger.Error(errNetBirdAPI, "Error updating Policy", "err", err)
nbPolicy.Status.Conditions = netbirdiov1.NBConditionFalse("APIError", fmt.Sprintf("Error updating policy: %v", err))
return policyID, false, err
}
requeue := false
if err != nil && strings.Contains(err.Error(), "not found") {
logger.Info("Policy deleted from NetBird API, recreating", "protocol", protocol)
policyID = nil
requeue = true
nbPolicy.Status.Conditions = netbirdiov1.NBConditionFalse("Gone", "Policy deleted from NetBird API")
} else if err != nil {
return nil, false, err
}
return policyID, requeue, nil
}
// Reconcile is part of the main kubernetes reconciliation loop which aims to
// move the current state of the cluster closer to the desired state.
func (r *NBPolicyReconciler) Reconcile(ctx context.Context, req ctrl.Request) (res ctrl.Result, err error) {
logger := ctrl.Log.WithName("NBPolicy").WithValues("namespace", req.Namespace, "name", req.Name)
logger.Info("Reconciling NBPolicy")
var nbPolicy netbirdiov1.NBPolicy
err = r.Client.Get(ctx, req.NamespacedName, &nbPolicy)
if err != nil {
if errors.IsNotFound(err) {
err = nil
}
if err != nil {
logger.Error(errKubernetesAPI, "error getting NBPolicy", "err", err)
}
return ctrl.Result{RequeueAfter: defaultRequeueAfter}, err
}
originalPolicy := nbPolicy.DeepCopy()
defer func() {
if err != nil {
// double check result is nil, otherwise error is not printed
// and exponential backoff doesn't work properly
res = ctrl.Result{}
return
}
if originalPolicy.DeletionTimestamp != nil && len(nbPolicy.Finalizers) == 0 {
return
}
if !originalPolicy.Status.Equal(nbPolicy.Status) {
updateErr := r.Client.Status().Update(ctx, &nbPolicy)
if updateErr != nil {
err = updateErr
}
}
if !res.Requeue && res.RequeueAfter == 0 {
res.RequeueAfter = defaultRequeueAfter
}
}()
if nbPolicy.DeletionTimestamp != nil {
if len(nbPolicy.Finalizers) == 0 {
return ctrl.Result{}, nil
}
return ctrl.Result{}, r.handleDelete(ctx, &nbPolicy, logger)
}
resourceList, err := r.getResources(ctx, &nbPolicy, logger)
if err != nil {
return ctrl.Result{}, err
}
portMapping, destGroups, err := r.mapResources(ctx, &nbPolicy, resourceList, logger)
if err != nil {
return ctrl.Result{}, err
}
sourceGroupIDs, err := r.groupNamesToIDs(ctx, nbPolicy.Spec.SourceGroups, logger)
if err != nil {
nbPolicy.Status.Conditions = netbirdiov1.NBConditionFalse("APIError", fmt.Sprintf("Error getting group IDs: %v", err))
return ctrl.Result{}, err
}
requeue, err := r.syncPolicy(ctx, &nbPolicy, sourceGroupIDs, destGroups, portMapping, logger)
if requeue || err != nil {
return ctrl.Result{Requeue: requeue}, err
}
nbPolicy.Status.Conditions = netbirdiov1.NBConditionTrue()
return ctrl.Result{}, nil
}
// syncPolicy ensure upstream policy is up-to-date
func (r *NBPolicyReconciler) syncPolicy(ctx context.Context, nbPolicy *netbirdiov1.NBPolicy, sourceGroups, destGroups []string, portMapping map[string][]int32, logger logr.Logger) (bool, error) {
requeue := false
for protocol, ports := range portMapping {
var policyID *string
switch protocol {
case protocolTCP:
policyID = nbPolicy.Status.TCPPolicyID
case protocolUDP:
policyID = nbPolicy.Status.UDPPolicyID
default:
logger.Error(errKubernetesAPI, "Unknown protocol", "protocol", protocol)
nbPolicy.Status.Conditions = netbirdiov1.NBConditionFalse("ConfigError", fmt.Sprintf("Unknown protocol: %s", protocol))
return requeue, errUnknownProtocol
}
if len(nbPolicy.Spec.Protocols) > 0 && !slices.Contains(nbPolicy.Spec.Protocols, protocol) {
if policyID != nil {
logger.Info("Deleting protocol policy as NBPolicy has restricted protocols", "protocol", protocol)
err := r.Netbird.Policies.Delete(ctx, *policyID)
if err != nil && !strings.Contains(err.Error(), "not found") {
nbPolicy.Status.Conditions = netbirdiov1.NBConditionFalse("APIError", fmt.Sprintf("Error deleting policy: %v", err))
return requeue, err
}
policyID = nil
} else {
logger.Info("Ignoring protocol as NBPolicy has restricted protocols", "protocol", protocol)
}
} else if len(ports) == 0 && policyID == nil {
logger.Info("0 ports found for protocol in policy", "protocol", protocol)
continue
} else if len(destGroups) == 0 && policyID == nil {
logger.Info("no destinations found for protocol in policy", "protocol", protocol)
continue
} else if len(sourceGroups) == 0 && policyID == nil {
logger.Info("no sources found for protocol in policy", "protocol", protocol)
continue
} else if len(ports) == 0 || len(destGroups) == 0 || len(sourceGroups) == 0 {
// Delete policy
logger.Info("Deleting policy", "protocol", protocol)
err := r.Netbird.Policies.Delete(ctx, *policyID)
if err != nil && !strings.Contains(err.Error(), "not found") {
nbPolicy.Status.Conditions = netbirdiov1.NBConditionFalse("APIError", fmt.Sprintf("Error deleting policy: %v", err))
return requeue, err
}
policyID = nil
} else {
var stringPorts []string
for _, v := range ports {
stringPorts = append(stringPorts, strconv.FormatInt(int64(v), 10))
}
for _, v := range nbPolicy.Spec.Ports {
stringPorts = append(stringPorts, strconv.FormatInt(int64(v), 10))
}
var err error
if policyID == nil {
policyID, err = r.createPolicy(ctx, nbPolicy, protocol, sourceGroups, destGroups, stringPorts, logger)
} else {
policyID, requeue, err = r.updatePolicy(ctx, policyID, nbPolicy, protocol, sourceGroups, destGroups, stringPorts, logger)
}
if err != nil {
return requeue, err
}
}
switch protocol {
case protocolTCP:
nbPolicy.Status.TCPPolicyID = policyID
case protocolUDP:
nbPolicy.Status.UDPPolicyID = policyID
default:
logger.Error(errKubernetesAPI, "Unknown protocol", "protocol", protocol)
nbPolicy.Status.Conditions = netbirdiov1.NBConditionFalse("ConfigError", fmt.Sprintf("Unknown protocol: %s", protocol))
return requeue, errUnknownProtocol
}
}
return requeue, nil
}
func (r *NBPolicyReconciler) handleDelete(ctx context.Context, nbPolicy *netbirdiov1.NBPolicy, logger logr.Logger) error {
if nbPolicy.Status.TCPPolicyID != nil {
err := r.Netbird.Policies.Delete(ctx, *nbPolicy.Status.TCPPolicyID)
if err != nil && !strings.Contains("not found", err.Error()) {
return err
}
nbPolicy.Status.TCPPolicyID = nil
}
if nbPolicy.Status.UDPPolicyID != nil {
err := r.Netbird.Policies.Delete(ctx, *nbPolicy.Status.UDPPolicyID)
if err != nil && !strings.Contains("not found", err.Error()) {
return err
}
nbPolicy.Status.UDPPolicyID = nil
}
if slices.Contains(nbPolicy.Finalizers, "netbird.io/cleanup") {
nbPolicy.Finalizers = util.Without(nbPolicy.Finalizers, "netbird.io/cleanup")
err := r.Client.Update(ctx, nbPolicy)
if err != nil {
logger.Error(errKubernetesAPI, "Error updating NBPolicy", "err", err)
return err
}
}
return nil
}
// groupNamesToIDs map list of NetBird group names to group IDs
func (r *NBPolicyReconciler) groupNamesToIDs(ctx context.Context, groupNames []string, logger logr.Logger) ([]string, error) {
groups, err := r.Netbird.Groups.List(ctx)
if err != nil {
logger.Error(errNetBirdAPI, "Error listing Groups", "err", err)
return nil, err
}
groupNameIDMapping := make(map[string]string)
for _, g := range groups {
groupNameIDMapping[g.Name] = g.Id
}
ret := make([]string, 0, len(groupNames))
for _, g := range groupNames {
ret = append(ret, groupNameIDMapping[g])
}
return ret, nil
}
// SetupWithManager sets up the controller with the Manager.
func (r *NBPolicyReconciler) SetupWithManager(mgr ctrl.Manager) error {
return ctrl.NewControllerManagedBy(mgr).
For(&netbirdiov1.NBPolicy{}).
Named("nbpolicy").
Complete(r)
}