From c274521faf40211b0daf3ab4da4b8f2ad6df25e8 Mon Sep 17 00:00:00 2001 From: Sammy Kerata Oina <44265300+SammyOina@users.noreply.github.com> Date: Mon, 13 May 2024 18:16:39 +0300 Subject: [PATCH] Refactor computation data handling to use filepaths (#126) Changed the internal representation of algorithms and datasets within the service from byte slices to file paths, writing received data directly to temp files. This modification allows for handling potentially large data sets without the need to load them entirely into memory, improving the memory efficiency and scalability of the service. Additionally, it aligns the call signature of external algorithms with the new approach, updating documentation and examples accordingly. Updated the linear regression example for consistency with the new data handling process. Resolves issues with memory bloat when processing large datasets. Signed-off-by: SammyOina --- agent/service.go | 74 ++++++++++++++++++++++--------------- test/manual/README.md | 2 +- test/manual/algo/lin_reg.py | 6 +-- 3 files changed, 49 insertions(+), 33 deletions(-) diff --git a/agent/service.go b/agent/service.go index cc5f0673..ec0dba20 100644 --- a/agent/service.go +++ b/agent/service.go @@ -64,8 +64,8 @@ type Service interface { type agentService struct { computation Computation // Holds the current computation request details. - algorithm []byte // Stores the algorithm received for the computation. - datasets [][]byte // Stores the datasets received for the computation. + algorithm string // Filepath to the algorithm received for the computation. + datasets []string // Filepath to the datasets received for the computation. result []byte // Stores the result of the computation. sm *StateMachine // Manages the state transitions of the agent service. runError error // Stores any error encountered during the computation run. @@ -100,7 +100,7 @@ func (as *agentService) Algo(ctx context.Context, algorithm Algorithm) error { if as.sm.GetState() != receivingAlgorithm { return errStateNotReady } - if as.algorithm != nil { + if as.algorithm != "" { return errAllManifestItemsReceived } @@ -118,9 +118,26 @@ func (as *agentService) Algo(ctx context.Context, algorithm Algorithm) error { return errHashMismatch } - as.algorithm = algorithm.Algorithm + f, err := os.CreateTemp("", "algorithm") + if err != nil { + return fmt.Errorf("error creating algorithm file: %v", err) + } - if as.algorithm != nil { + if _, err := f.Write(algorithm.Algorithm); err != nil { + return fmt.Errorf("error writing algorithm to file: %v", err) + } + + if err := os.Chmod(f.Name(), algoFilePermission); err != nil { + return fmt.Errorf("error changing file permissions: %v", err) + } + + if err := f.Close(); err != nil { + return fmt.Errorf("error closing file: %v", err) + } + + as.algorithm = f.Name() + + if as.algorithm != "" { as.sm.SendEvent(algorithmReceived) } @@ -151,7 +168,19 @@ func (as *agentService) Data(ctx context.Context, dataset Dataset) error { as.computation.Datasets = slices.Delete(as.computation.Datasets, index, index+1) } - as.datasets = append(as.datasets, dataset.Dataset) + f, err := os.CreateTemp("", fmt.Sprintf("dataset-%s", dataset.ID)) + if err != nil { + return fmt.Errorf("error creating dataset file: %v", err) + } + + if _, err := f.Write(dataset.Dataset); err != nil { + return fmt.Errorf("error writing dataset to file: %v", err) + } + if err := f.Close(); err != nil { + return fmt.Errorf("error closing file: %v", err) + } + + as.datasets = append(as.datasets, f.Name()) if len(as.computation.Datasets) == 0 { as.sm.SendEvent(dataReceived) @@ -200,7 +229,7 @@ func (as *agentService) runComputation() { as.sm.logger.Debug("computation run started") defer as.sm.SendEvent(runComplete) as.publishEvent("in-progress", json.RawMessage{})() - result, err := run(as.algorithm, as.datasets[0]) + result, err := run(as.algorithm, as.datasets) if err != nil { as.runError = err as.sm.logger.Warn(fmt.Sprintf("computation failed with error: %s", err.Error())) @@ -219,7 +248,13 @@ func (as *agentService) publishEvent(status string, details json.RawMessage) fun } } -func run(algoContent, dataContent []byte) ([]byte, error) { +func run(algoFile string, dataFiles []string) ([]byte, error) { + defer os.Remove(algoFile) + defer func() { + for _, file := range dataFiles { + os.Remove(file) + } + }() listener, err := socket.StartUnixSocketServer(socketPath) if err != nil { return nil, fmt.Errorf("error creating stdout pipe: %v", err) @@ -234,27 +269,8 @@ func run(algoContent, dataContent []byte) ([]byte, error) { go socket.AcceptConnection(listener, dataChannel, errorChannel) - f, err := os.CreateTemp("", "algorithm") - if err != nil { - return nil, fmt.Errorf("error creating algorithm file: %v", err) - } - defer os.Remove(f.Name()) - - if _, err := f.Write(algoContent); err != nil { - return nil, fmt.Errorf("error writing algorithm to file: %v", err) - } - - if err := os.Chmod(f.Name(), algoFilePermission); err != nil { - return nil, fmt.Errorf("error changing file permissions: %v", err) - } - - if err := f.Close(); err != nil { - return nil, fmt.Errorf("error closing file: %v", err) - } - - // Construct the executable with CSV data as a command-line argument - data := string(dataContent) - cmd := exec.Command(f.Name(), data, socketPath) + args := append([]string{socketPath}, dataFiles...) + cmd := exec.Command(algoFile, args...) if err := cmd.Start(); err != nil { return nil, fmt.Errorf("error starting algorithm: %v", err) diff --git a/test/manual/README.md b/test/manual/README.md index 3adcf578..945b9305 100644 --- a/test/manual/README.md +++ b/test/manual/README.md @@ -8,7 +8,7 @@ Throughout the tests, we assume that our current working directory is the root o Agent accepts the algorithm as a binary that take in two command line arguments. ```shell -algorithm-file +algorithm-file ``` The algorithm program should return the results to a socket and an example can be seen in this [file](./algo/lin_reg.py). diff --git a/test/manual/algo/lin_reg.py b/test/manual/algo/lin_reg.py index 92f0ffed..191acaa4 100644 --- a/test/manual/algo/lin_reg.py +++ b/test/manual/algo/lin_reg.py @@ -6,8 +6,8 @@ import pandas as pd from sklearn.model_selection import train_test_split from sklearn.linear_model import LogisticRegression -dataset = sys.argv[1] -iris = pd.read_csv(io.StringIO(dataset)) +csv_file_path = sys.argv[2] +iris = pd.read_csv(csv_file_path) # Droping the Species since we only need the measurements X = iris.drop(['Species'], axis=1) @@ -30,7 +30,7 @@ joblib.dump(log_reg, model_buffer) model_bytes = model_buffer.getvalue() # Define the path for the Unix domain socket -socket_path = sys.argv[2] +socket_path = sys.argv[1] # Create a Unix domain socket client client = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM)