Skip to content

Commit 49a3955

Browse files
yanmxaclaude
andcommitted
refactor: simplify OpenFL-only code paths in server and client files
Remove unnecessary framework switch statements and embed.FS variables since federatedLearningServer() and clusterWorkload() are now only called for the OpenFL path. Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> Signed-off-by: Meng Yan <myan@redhat.com>
1 parent 9fad429 commit 49a3955

2 files changed

Lines changed: 53 additions & 70 deletions

File tree

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

Lines changed: 26 additions & 35 deletions
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,6 @@ package controller
22

33
import (
44
"context"
5-
"embed"
65
"fmt"
76
"net"
87
"reflect"
@@ -334,43 +333,35 @@ func (r *FederatedLearningReconciler) clusterWorkload(ctx context.Context, insta
334333
obsSidecarImage = instance.ObjectMeta.Annotations[v1alpha1.AnnotationSidecarImage]
335334
}
336335

337-
var clientParams any
338-
var clientFS embed.FS
336+
host, port, err := net.SplitHostPort(serverAddress)
337+
if err != nil {
338+
return fmt.Errorf("failed to parse server address: %w", err)
339+
}
340+
portUint, err := strconv.ParseUint(port, 10, 16)
341+
if err != nil {
342+
return fmt.Errorf("failed to parse server port: %w", err)
343+
}
344+
modelDir, _, err := getDirFile(instance.Spec.Server.Storage.ModelPath)
345+
if err != nil {
346+
return err
347+
}
339348

340-
switch instance.Spec.Framework {
341-
case flv1alpha1.OpenFL:
342-
clientFS = manifests.OpenFLClientFiles
343-
host, port, err := net.SplitHostPort(serverAddress)
344-
if err != nil {
345-
return fmt.Errorf("failed to parse server address: %w", err)
346-
}
347-
portUint, err := strconv.ParseUint(port, 10, 16)
348-
if err != nil {
349-
return fmt.Errorf("failed to parse server port: %w", err)
350-
}
351-
modelDir, _, err := getDirFile(instance.Spec.Server.Storage.ModelPath)
352-
if err != nil {
353-
return err
354-
}
355-
clientParams = &manifests.OpenFLClientParams{
356-
ManifestName: instance.Name,
357-
ManifestNamespace: clusterName,
358-
ClientJobNamespace: instance.Namespace,
359-
ClientJobName: fmt.Sprintf("%s-client", instance.Name),
360-
ClientJobImage: instance.Spec.Client.Image,
361-
ClientDataPath: dataConfig,
362-
ServerIP: host,
363-
ServerPort: uint16(portUint),
364-
ModelDir: modelDir,
365-
ObsSidecarImage: obsSidecarImage,
366-
ClientName: clusterName,
367-
NumberOfRounds: instance.Spec.Server.Rounds,
368-
}
369-
default:
370-
return fmt.Errorf("unsupported framework: %s", instance.Spec.Framework)
349+
clientParams := &manifests.OpenFLClientParams{
350+
ManifestName: instance.Name,
351+
ManifestNamespace: clusterName,
352+
ClientJobNamespace: instance.Namespace,
353+
ClientJobName: fmt.Sprintf("%s-client", instance.Name),
354+
ClientJobImage: instance.Spec.Client.Image,
355+
ClientDataPath: dataConfig,
356+
ServerIP: host,
357+
ServerPort: uint16(portUint),
358+
ModelDir: modelDir,
359+
ObsSidecarImage: obsSidecarImage,
360+
ClientName: clusterName,
361+
NumberOfRounds: instance.Spec.Server.Rounds,
371362
}
372363

373-
render, deployer := applier.NewRenderer(clientFS), applier.NewDeployer(r.Client)
364+
render, deployer := applier.NewRenderer(manifests.OpenFLClientFiles), applier.NewDeployer(r.Client)
374365
unstructuredObjects, err := render.Render("", "", func(profile string) (interface{}, error) {
375366
return clientParams, nil
376367
})

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

Lines changed: 27 additions & 35 deletions
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,6 @@ package controller
22

33
import (
44
"context"
5-
"embed"
65
"fmt"
76
"sort"
87
"strings"
@@ -120,43 +119,36 @@ func (r *FederatedLearningReconciler) federatedLearningServer(ctx context.Contex
120119
obsSidecarImage = instance.ObjectMeta.Annotations[v1alpha1.AnnotationSidecarImage]
121120
}
122121

123-
var serverParams any
124-
var serverFS embed.FS
122+
clusters, err := r.getDecidedClusters(ctx, instance)
123+
if err != nil {
124+
return err
125+
}
126+
sort.Strings(clusters)
127+
log.Infof("clusters: %+v", clusters)
125128

126-
switch instance.Spec.Framework {
127-
case flv1alpha1.OpenFL:
128-
serverFS = manifests.OpenFLServerFiles
129-
clusters, err := r.getDecidedClusters(ctx, instance)
130-
if err != nil {
131-
return err
132-
}
133-
sort.Strings(clusters)
134-
log.Infof("clusters: %+v", clusters)
135-
// determine endpoint info (IP and port) prior to rendering, especially for NodePort
136-
listenerIP, listenerPort, err := r.determineEndpointInfo(ctx, instance)
137-
if err != nil {
138-
return err
139-
}
140-
serverParams = &manifests.OpenFLServerParams{
141-
Namespace: instance.Namespace,
142-
Name: getSeverName(instance.Name),
143-
Image: instance.Spec.Server.Image,
144-
NumberOfRounds: instance.Spec.Server.Rounds,
145-
StorageVolumeName: instance.Spec.Server.Storage.Name,
146-
ListenerType: string(instance.Spec.Server.Listeners[0].Type),
147-
ListenerIP: listenerIP,
148-
ListenerPort: listenerPort,
149-
CreateService: createService,
150-
ModelDir: modelDir,
151-
ObsSidecarImage: obsSidecarImage,
152-
Collaborators: strings.Join(clusters, ","),
153-
}
154-
log.Infof("server params: %+v", serverParams)
155-
default:
156-
return fmt.Errorf("unsupported framework: %s", instance.Spec.Framework)
129+
// determine endpoint info (IP and port) prior to rendering, especially for NodePort
130+
listenerIP, listenerPort, err := r.determineEndpointInfo(ctx, instance)
131+
if err != nil {
132+
return err
133+
}
134+
135+
serverParams := &manifests.OpenFLServerParams{
136+
Namespace: instance.Namespace,
137+
Name: getSeverName(instance.Name),
138+
Image: instance.Spec.Server.Image,
139+
NumberOfRounds: instance.Spec.Server.Rounds,
140+
StorageVolumeName: instance.Spec.Server.Storage.Name,
141+
ListenerType: string(instance.Spec.Server.Listeners[0].Type),
142+
ListenerIP: listenerIP,
143+
ListenerPort: listenerPort,
144+
CreateService: createService,
145+
ModelDir: modelDir,
146+
ObsSidecarImage: obsSidecarImage,
147+
Collaborators: strings.Join(clusters, ","),
157148
}
149+
log.Infof("server params: %+v", serverParams)
158150

159-
render, deployer := applier.NewRenderer(serverFS), applier.NewDeployer(r.Client)
151+
render, deployer := applier.NewRenderer(manifests.OpenFLServerFiles), applier.NewDeployer(r.Client)
160152
unstructuredObjects, err := render.Render("", "", func(profile string) (interface{}, error) {
161153
return serverParams, nil
162154
})

0 commit comments

Comments
 (0)