Files
netbird-kubernetes-operator/internal/controller/nbpolicy_controller.go
T
Philip LaineandGitHub 60ce12c74e Update Netbird dependency and Golang version (#112)
This change updates the Netbird dependency to the latest version to
enabled #111 to be implemented with future API additions. This change
requires updating the Go version as the upstream dependency requires it.

An interesting aspect that was required was to pin the dex dependency to
v2. It seems like the Netbird go.mod is doing something unexpected with
their version.

Signed-off-by: Philip Laine <philip.laine@gmail.com>
2026-03-03 10:31:38 +01:00

408 lines
14 KiB
Go

package controller
import (
"context"
"fmt"
"slices"
"strconv"
"strings"
"k8s.io/apimachinery/pkg/api/errors"
"k8s.io/apimachinery/pkg/runtime"
"k8s.io/apimachinery/pkg/types"
ctrl "sigs.k8s.io/controller-runtime"
"sigs.k8s.io/controller-runtime/pkg/client"
"github.com/go-logr/logr"
netbirdiov1 "github.com/netbirdio/kubernetes-operator/api/v1"
"github.com/netbirdio/kubernetes-operator/internal/util"
netbird "github.com/netbirdio/netbird/shared/management/client/rest"
"github.com/netbirdio/netbird/shared/management/http/api"
)
// NBPolicyReconciler reconciles a NBPolicy object
type NBPolicyReconciler struct {
client.Client
Scheme *runtime.Scheme
ClusterName string
APIKey string
ManagementURL string
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]interface{}{
protocolTCP: make(map[int32]interface{}),
protocolUDP: make(map[int32]interface{}),
}
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 == "" && !util.Contains(resourcePolicies, nbPolicy.Name) {
continue
}
if generatedBy != "" && !util.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 && !util.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 util.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 {
r.netbird = netbird.New(r.ManagementURL, r.APIKey)
return ctrl.NewControllerManagedBy(mgr).
For(&netbirdiov1.NBPolicy{}).
Named("nbpolicy").
Complete(r)
}