Skip to content

Commit 0e29935

Browse files
authored
✨ Upload the model to AWS S3 in federated learning controller (open-cluster-management-io#80)
* add S3 api in types Signed-off-by: mrrr61 <mrrr61@outlook.com> * feat(crd): add s3 storage configuration support Signed-off-by: mrrr61 <mrrr61@outlook.com> * mount aws s3 pvc and upload the model in Flower framework Signed-off-by: mrrr61 <mrrr61@outlook.com> * chore: regenerate FL CRD and deepcopy artifacts Signed-off-by: mrrr61 <mrrr61@outlook.com> * refactor(federated-learning-controller): rename storage types for clarity Signed-off-by: mrrr61 <mrrr61@outlook.com> * fix(api): correct S3Bucket storage type value Signed-off-by: mrrr61 <mrrr61@outlook.com> --------- Signed-off-by: mrrr61 <mrrr61@outlook.com>
1 parent f5f661d commit 0e29935

4 files changed

Lines changed: 278 additions & 14 deletions

File tree

federated-learning-controller/api/v1alpha1/federatedlearning_types.go

Lines changed: 25 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -86,18 +86,37 @@ type ServerSpec struct {
8686

8787
// ModelStorageSpec defines the storage specification for the model.
8888
type ModelStorageSpec struct {
89-
Name string `json:"name,omitempty"`
90-
Type StorageType `json:"type,omitempty"`
91-
Size string `json:"size,omitempty"` // +optional
92-
ModelPath string `json:"path,omitempty"` //
89+
Name string `json:"name,omitempty"`
90+
Type StorageType `json:"type,omitempty"`
91+
Size string `json:"size,omitempty"` // +optional
92+
ModelPath string `json:"path,omitempty"` //
93+
S3 *S3StorageSpec `json:"s3,omitempty"`
9394
}
9495

9596
// StorageType represents the type of storage.
9697
type StorageType string
9798

9899
const (
99-
PersistentVolumeClaim StorageType = "PersistentVolumeClaim"
100-
HostPathStorage StorageType = "HostPath"
100+
S3Driver = "s3.csi.aws.com"
101+
S3VolumeHandle = "s3-csi-driver-volume"
102+
)
103+
104+
// S3StorageSpec defines the configuration required for S3-backed persistent volumes.
105+
type S3StorageSpec struct {
106+
// +kubebuilder:validation:Required
107+
BucketName string `json:"bucketName"`
108+
// Region configures the AWS region for the S3 connection.
109+
// +optional
110+
Region string `json:"region,omitempty"`
111+
// Prefix specifies the object prefix (folder) inside the bucket to use.
112+
// +optional
113+
Prefix string `json:"prefix,omitempty"`
114+
}
115+
116+
const (
117+
PVCStorage StorageType = "PersistentVolumeClaim"
118+
HostPath StorageType = "HostPath"
119+
S3Bucket StorageType = "S3Bucket"
101120
)
102121

103122
// ListenerSpec defines the specification for a listener.

federated-learning-controller/api/v1alpha1/zz_generated.deepcopy.go

Lines changed: 21 additions & 1 deletion
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

federated-learning-controller/config/crd/bases/federation-ai.open-cluster-management.io_federatedlearnings.yaml

Lines changed: 17 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -573,6 +573,23 @@ spec:
573573
type: string
574574
path:
575575
type: string
576+
s3:
577+
description: S3StorageSpec defines the configuration required
578+
for S3-backed persistent volumes.
579+
properties:
580+
bucketName:
581+
type: string
582+
prefix:
583+
description: Prefix specifies the object prefix (folder)
584+
inside the bucket to use.
585+
type: string
586+
region:
587+
description: Region configures the AWS region for the
588+
S3 connection.
589+
type: string
590+
required:
591+
- bucketName
592+
type: object
576593
size:
577594
type: string
578595
type:

federated-learning-controller/internal/controller/federatedlearning_server.go

Lines changed: 215 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -612,19 +612,34 @@ func SetOwner(objects []*unstructured.Unstructured,
612612
}
613613

614614
func (r *FederatedLearningReconciler) storage(ctx context.Context, instance *flv1alpha1.FederatedLearning) error {
615-
namespace := instance.Namespace
616-
name := instance.Spec.Server.Storage.Name
617-
size := instance.Spec.Server.Storage.Size
618615
storageType := instance.Spec.Server.Storage.Type
619-
if storageType != flv1alpha1.PersistentVolumeClaim {
616+
617+
switch storageType {
618+
case flv1alpha1.PVCStorage:
619+
return r.ensureStandardPVC(ctx, instance)
620+
case flv1alpha1.S3Bucket:
621+
return r.ensureS3PVC(ctx, instance)
622+
default:
620623
return fmt.Errorf("unsupported storage type: %s", storageType)
621624
}
625+
}
626+
627+
func (r *FederatedLearningReconciler) ensureStandardPVC(ctx context.Context, instance *flv1alpha1.FederatedLearning) error {
628+
namespace := instance.Namespace
629+
name := instance.Spec.Server.Storage.Name
630+
631+
if instance.Spec.Server.Storage.Size == "" {
632+
return fmt.Errorf("size must be specified for storage type %s", instance.Spec.Server.Storage.Type)
633+
}
622634

623635
pvc := &corev1.PersistentVolumeClaim{}
624636
err := r.Get(ctx, client.ObjectKey{Namespace: namespace, Name: name}, pvc)
625637
if err != nil {
626638
if errors.IsNotFound(err) {
627-
// PVC does not exist, create it
639+
quantity, parseErr := resource.ParseQuantity(instance.Spec.Server.Storage.Size)
640+
if parseErr != nil {
641+
return fmt.Errorf("failed to parse storage size %q: %w", instance.Spec.Server.Storage.Size, parseErr)
642+
}
628643
newPVC := &corev1.PersistentVolumeClaim{
629644
ObjectMeta: metav1.ObjectMeta{
630645
Name: name,
@@ -636,7 +651,7 @@ func (r *FederatedLearningReconciler) storage(ctx context.Context, instance *flv
636651
},
637652
Resources: corev1.VolumeResourceRequirements{
638653
Requests: corev1.ResourceList{
639-
corev1.ResourceStorage: resource.MustParse(size),
654+
corev1.ResourceStorage: quantity,
640655
},
641656
},
642657
},
@@ -650,7 +665,200 @@ func (r *FederatedLearningReconciler) storage(ctx context.Context, instance *flv
650665
return err
651666
}
652667

653-
// PVC exists
654668
log.Infof("storage PVC already exists: %s", name)
655669
return nil
656670
}
671+
672+
func (r *FederatedLearningReconciler) ensureS3PVC(ctx context.Context, instance *flv1alpha1.FederatedLearning) error {
673+
storageSpec := instance.Spec.Server.Storage
674+
namespace := instance.Namespace
675+
claimName := storageSpec.Name
676+
677+
if storageSpec.Size == "" {
678+
return fmt.Errorf("size must be specified for storage type %s", storageSpec.Type)
679+
}
680+
if storageSpec.S3 == nil {
681+
return fmt.Errorf("s3 configuration must be provided when storage type is %s", storageSpec.Type)
682+
}
683+
if storageSpec.S3.BucketName == "" {
684+
return fmt.Errorf("bucketName is required for s3 storage")
685+
}
686+
687+
requestQuantity, err := resource.ParseQuantity(storageSpec.Size)
688+
if err != nil {
689+
return fmt.Errorf("failed to parse storage size %q: %w", storageSpec.Size, err)
690+
}
691+
pvName := fmt.Sprintf("%s-pv", claimName)
692+
693+
mountOptions := make([]string, 0)
694+
if storageSpec.S3.Region != "" {
695+
mountOptions = append(mountOptions, fmt.Sprintf("region %s", storageSpec.S3.Region))
696+
}
697+
if storageSpec.S3.Prefix != "" {
698+
mountOptions = append(mountOptions, fmt.Sprintf("prefix %s", storageSpec.S3.Prefix))
699+
}
700+
desiredPV := &corev1.PersistentVolume{
701+
ObjectMeta: metav1.ObjectMeta{
702+
Name: pvName,
703+
Namespace: namespace,
704+
},
705+
Spec: corev1.PersistentVolumeSpec{
706+
AccessModes: []corev1.PersistentVolumeAccessMode{
707+
corev1.ReadWriteMany,
708+
},
709+
Capacity: corev1.ResourceList{
710+
corev1.ResourceStorage: requestQuantity,
711+
},
712+
ClaimRef: &corev1.ObjectReference{
713+
Namespace: namespace,
714+
Name: claimName,
715+
},
716+
PersistentVolumeSource: corev1.PersistentVolumeSource{
717+
CSI: &corev1.CSIPersistentVolumeSource{
718+
Driver: v1alpha1.S3Driver,
719+
VolumeHandle: v1alpha1.S3VolumeHandle,
720+
VolumeAttributes: map[string]string{
721+
"bucketName": storageSpec.S3.BucketName,
722+
},
723+
},
724+
},
725+
MountOptions: mountOptions,
726+
StorageClassName: "",
727+
},
728+
}
729+
existingPV := &corev1.PersistentVolume{}
730+
if err := r.Get(ctx, client.ObjectKey{Name: pvName}, existingPV); err != nil {
731+
if !errors.IsNotFound(err) {
732+
return err
733+
}
734+
if err := r.Create(ctx, desiredPV); err != nil {
735+
return fmt.Errorf("failed to create s3 persistent volume %q: %w", pvName, err)
736+
}
737+
log.Infow("created S3 persistent volume", "name", pvName)
738+
} else {
739+
updated := false
740+
if existingPV.Spec.PersistentVolumeSource.CSI == nil {
741+
existingPV.Spec.PersistentVolumeSource.CSI = &corev1.CSIPersistentVolumeSource{}
742+
updated = true
743+
}
744+
csiSpec := existingPV.Spec.PersistentVolumeSource.CSI
745+
if csiSpec.Driver != v1alpha1.S3Driver {
746+
csiSpec.Driver = v1alpha1.S3Driver
747+
updated = true
748+
}
749+
if csiSpec.VolumeHandle != v1alpha1.S3VolumeHandle {
750+
csiSpec.VolumeHandle = v1alpha1.S3VolumeHandle
751+
updated = true
752+
}
753+
if csiSpec.VolumeAttributes == nil {
754+
csiSpec.VolumeAttributes = map[string]string{}
755+
}
756+
if csiSpec.VolumeAttributes["bucketName"] != storageSpec.S3.BucketName {
757+
csiSpec.VolumeAttributes["bucketName"] = storageSpec.S3.BucketName
758+
updated = true
759+
}
760+
if len(existingPV.Spec.AccessModes) != 1 || existingPV.Spec.AccessModes[0] != corev1.ReadWriteMany {
761+
existingPV.Spec.AccessModes = []corev1.PersistentVolumeAccessMode{corev1.ReadWriteMany}
762+
updated = true
763+
}
764+
if qty, ok := existingPV.Spec.Capacity[corev1.ResourceStorage]; !ok || qty.Cmp(requestQuantity) != 0 {
765+
existingPV.Spec.Capacity[corev1.ResourceStorage] = requestQuantity
766+
updated = true
767+
}
768+
if existingPV.Spec.ClaimRef == nil ||
769+
existingPV.Spec.ClaimRef.Namespace != namespace ||
770+
existingPV.Spec.ClaimRef.Name != claimName {
771+
existingPV.Spec.ClaimRef = &corev1.ObjectReference{
772+
Namespace: namespace,
773+
Name: claimName,
774+
}
775+
updated = true
776+
}
777+
if existingPV.Spec.StorageClassName != "" {
778+
existingPV.Spec.StorageClassName = ""
779+
updated = true
780+
}
781+
if !equalStringSlice(existingPV.Spec.MountOptions, mountOptions) {
782+
existingPV.Spec.MountOptions = mountOptions
783+
updated = true
784+
}
785+
if updated {
786+
if err := r.Update(ctx, existingPV); err != nil {
787+
return fmt.Errorf("failed to update s3 persistent volume %q: %w", pvName, err)
788+
}
789+
log.Infow("updated S3 persistent volume", "name", pvName)
790+
}
791+
}
792+
793+
desiredPVC := &corev1.PersistentVolumeClaim{
794+
ObjectMeta: metav1.ObjectMeta{
795+
Name: claimName,
796+
Namespace: namespace,
797+
},
798+
Spec: corev1.PersistentVolumeClaimSpec{
799+
AccessModes: []corev1.PersistentVolumeAccessMode{
800+
corev1.ReadWriteMany,
801+
},
802+
Resources: corev1.VolumeResourceRequirements{
803+
Requests: corev1.ResourceList{
804+
corev1.ResourceStorage: requestQuantity,
805+
},
806+
},
807+
VolumeName: pvName,
808+
},
809+
}
810+
storageClass := ""
811+
desiredPVC.Spec.StorageClassName = &storageClass
812+
813+
existingPVC := &corev1.PersistentVolumeClaim{}
814+
if err := r.Get(ctx, client.ObjectKey{Namespace: namespace, Name: claimName}, existingPVC); err != nil {
815+
if !errors.IsNotFound(err) {
816+
return err
817+
}
818+
if err := r.Create(ctx, desiredPVC); err != nil {
819+
return fmt.Errorf("failed to create s3 persistent volume claim %q: %w", claimName, err)
820+
}
821+
log.Infow("created S3 persistent volume claim", "name", claimName, "namespace", namespace)
822+
return nil
823+
}
824+
updated := false
825+
if len(existingPVC.Spec.AccessModes) != 1 || existingPVC.Spec.AccessModes[0] != corev1.ReadWriteMany {
826+
existingPVC.Spec.AccessModes = []corev1.PersistentVolumeAccessMode{corev1.ReadWriteMany}
827+
updated = true
828+
}
829+
if existingPVC.Spec.StorageClassName == nil || *existingPVC.Spec.StorageClassName != "" {
830+
existingPVC.Spec.StorageClassName = &storageClass
831+
updated = true
832+
}
833+
if existingPVC.Spec.VolumeName != pvName {
834+
existingPVC.Spec.VolumeName = pvName
835+
updated = true
836+
}
837+
if existingPVC.Spec.Resources.Requests == nil {
838+
existingPVC.Spec.Resources.Requests = corev1.ResourceList{}
839+
}
840+
if qty, ok := existingPVC.Spec.Resources.Requests[corev1.ResourceStorage]; !ok || qty.Cmp(requestQuantity) != 0 {
841+
existingPVC.Spec.Resources.Requests[corev1.ResourceStorage] = requestQuantity
842+
updated = true
843+
}
844+
if updated {
845+
if err := r.Update(ctx, existingPVC); err != nil {
846+
return fmt.Errorf("failed to update s3 persistent volume claim %q: %w", claimName, err)
847+
}
848+
log.Infow("updated S3 persistent volume claim", "name", claimName, "namespace", namespace)
849+
return nil
850+
}
851+
log.Infof("S3 PVC already configured: %s/%s", namespace, claimName)
852+
return nil
853+
}
854+
func equalStringSlice(a, b []string) bool {
855+
if len(a) != len(b) {
856+
return false
857+
}
858+
for i := range a {
859+
if a[i] != b[i] {
860+
return false
861+
}
862+
}
863+
return true
864+
}

0 commit comments

Comments
 (0)