mirror of
https://github.com/ultravioletrs/cocos.git
synced 2026-08-07 07:14:50 +00:00
NOISSUE - Enhance OCI image extraction to return algorithm and requirements paths, and add deferred cleanup for temporary files (#586)
CI / lint (push) Has been cancelled
CI / test (agent) (push) Has been cancelled
CI / test (cli) (push) Has been cancelled
CI / test (cmd) (push) Has been cancelled
CI / test (internal) (push) Has been cancelled
CI / test (manager, true) (push) Has been cancelled
CI / test (pkg) (push) Has been cancelled
CI / upload-coverage (push) Has been cancelled
CI / lint (push) Has been cancelled
CI / test (agent) (push) Has been cancelled
CI / test (cli) (push) Has been cancelled
CI / test (cmd) (push) Has been cancelled
CI / test (internal) (push) Has been cancelled
CI / test (manager, true) (push) Has been cancelled
CI / test (pkg) (push) Has been cancelled
CI / upload-coverage (push) Has been cancelled
* feat: Enhance OCI image extraction to return algorithm and requirements paths, and add deferred cleanup for temporary files. Signed-off-by: Sammy Oina <sammyoina@gmail.com> * feat: implement deterministic zipping and enhance checksum verification for resources Signed-off-by: Sammy Oina <sammyoina@gmail.com> * feat: Update component build sources, add gRPC health checks to the CVM server, and refine algorithm argument handling and documentation. Signed-off-by: Sammy Oina <sammyoina@gmail.com> * docs: Update remote resources testing guide with `sudo` for KBS, algorithm result saving, `requirements.txt`, and `algo-args` for RVPS. Signed-off-by: Sammy Oina <sammyoina@gmail.com> * refactor: Explicitly ignore `stderr.Write` return values and add minor whitespace in tests. Signed-off-by: Sammy Oina <sammyoina@gmail.com> * test: add comprehensive error path and edge case tests for file, zip, OCI, and agent components. Signed-off-by: Sammy Oina <sammyoina@gmail.com> * feat: Add mutexes for thread-safe algorithm execution and expand recognized data file extensions to include common archive formats. Signed-off-by: Sammy Oina <sammyoina@gmail.com> * feat: Add OCI extraction tests for Python algorithms and multi-layer datasets, refactor algorithm execution for testability, and enhance algorithm stop and error handling tests. Signed-off-by: Sammy Oina <sammyoina@gmail.com> * test: Add error assertions to OCI extraction test helpers and remove an unused mock exec command. Signed-off-by: Sammy Oina <sammyoina@gmail.com> * test: Improve error handling test coverage for algorithm execution and OCI resource extraction. Signed-off-by: Sammy Oina <sammyoina@gmail.com> * fix: Improve algorithm process termination, enhance computation error handling, and add concurrency safety to agent service. Signed-off-by: Sammy Oina <sammyoina@gmail.com> --------- Signed-off-by: Sammy Oina <sammyoina@gmail.com>
This commit is contained in:
committed by
GitHub
parent
80bf813c48
commit
b44780df95
@@ -92,7 +92,7 @@ EOF
|
||||
mkdir -p kbs-data/as kbs-data/rvps kbs-data/repository
|
||||
|
||||
# Start KBS
|
||||
../target/release/kbs --config-file kbs-config.toml
|
||||
sudo ../target/release/kbs --config-file kbs-config.toml
|
||||
```
|
||||
|
||||
KBS will listen on `http://localhost:8080`
|
||||
@@ -115,6 +115,7 @@ cat > lin_reg.py << 'EOF'
|
||||
import pandas as pd
|
||||
from sklearn.linear_model import LinearRegression
|
||||
import sys
|
||||
import os
|
||||
|
||||
# Load dataset
|
||||
data = pd.read_csv(sys.argv[1])
|
||||
@@ -126,34 +127,46 @@ model = LinearRegression()
|
||||
model.fit(X, y)
|
||||
|
||||
# Save results
|
||||
os.makedirs("results", exist_ok=True)
|
||||
with open("results/output.txt", "w") as f:
|
||||
f.write(f"Coefficients: {model.coef_}\n")
|
||||
f.write(f"Intercept: {model.intercept_}\n")
|
||||
|
||||
print(f"Coefficients: {model.coef_}")
|
||||
print(f"Intercept: {model.intercept_}")
|
||||
EOF
|
||||
|
||||
# 2. Create a Dockerfile
|
||||
# 2. Create requirements.txt
|
||||
cat > requirements.txt << 'EOF'
|
||||
pandas
|
||||
scikit-learn
|
||||
EOF
|
||||
|
||||
# 3. Create a Dockerfile
|
||||
cat > Dockerfile << 'EOF'
|
||||
FROM python:3.9-slim
|
||||
RUN pip install pandas scikit-learn
|
||||
COPY lin_reg.py /app/algorithm.py
|
||||
COPY requirements.txt /app/requirements.txt
|
||||
WORKDIR /app
|
||||
ENTRYPOINT ["python", "algorithm.py"]
|
||||
EOF
|
||||
|
||||
# 3. Build the image
|
||||
# 4. Build the image
|
||||
docker build -t localhost:5000/lin-reg-algo:v1.0 .
|
||||
docker push localhost:5000/lin-reg-algo:v1.0
|
||||
|
||||
# 4. Generate and store key
|
||||
# 5. Generate and store key
|
||||
openssl rand -out algo.key 32
|
||||
|
||||
# 5. Store key in KBS using kbs-client
|
||||
# 6. Store key in KBS using kbs-client
|
||||
../target/release/kbs-client --url http://localhost:8080 config \
|
||||
--auth-private-key kbs-admin.key \
|
||||
set-resource \
|
||||
--path default/key/algo-key \
|
||||
--resource-file algo.key
|
||||
|
||||
# 6. Encrypt the image using Host Skopeo + Docker Keyprovider
|
||||
# 7. Encrypt the image using Host Skopeo + Docker Keyprovider
|
||||
# Start Keyprovider in background
|
||||
docker run -d --rm --name keyprovider --network host \
|
||||
-v "$PWD:/work" -w /work \
|
||||
@@ -255,10 +268,12 @@ HOST_IP=$(ip -4 addr show | grep -oP '(?<=inet\s)\d+(\.\d+){3}' | grep -v 127.0.
|
||||
|
||||
Start CVMS server:
|
||||
```bash
|
||||
# Calculate SHA3-256 of decrypted files using cocos-cli
|
||||
# Calculate SHA3-256 of decrypted files using cocos-cli or cvms-test
|
||||
# NOTE: We use the hash of the original plaintext files, as the Agent validates the decrypted content.
|
||||
# Redirect stderr to stdout (2>&1) because cocos-cli prints to stderr
|
||||
# For single files, use the file hash. For directories, use the hash of the directory (which the tools zip deterministically).
|
||||
|
||||
ALGO_HASH=$(./build/cocos-cli checksum lin_reg.py 2>&1 | awk '{print $NF}')
|
||||
|
||||
DATASET_HASH=$(./build/cocos-cli checksum iris.csv 2>&1 | awk '{print $NF}')
|
||||
|
||||
go build -o build/cvms-test ./test/cvms/main.go
|
||||
@@ -266,11 +281,11 @@ HOST=$HOST_IP PORT=7001 ./build/cvms-test \
|
||||
-public-key-path ./public.pem \
|
||||
-attested-tls-bool false \
|
||||
-kbs-url http://$HOST_IP:8080 \
|
||||
-algo-type oci-image \
|
||||
-algo-type python \
|
||||
-algo-source-url docker://$HOST_IP:5000/encrypted-lin-reg:v1.0 \
|
||||
-algo-kbs-path default/key/algo-key \
|
||||
-algo-hash $ALGO_HASH \
|
||||
-dataset-type oci-image \
|
||||
-algo-args datasets/dataset_0.csv \
|
||||
-dataset-source-urls docker://$HOST_IP:5000/encrypted-iris:v1.0 \
|
||||
-dataset-kbs-paths default/key/dataset-key \
|
||||
-dataset-hash $DATASET_HASH
|
||||
|
||||
@@ -3,16 +3,21 @@
|
||||
package binary
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"os"
|
||||
"os/exec"
|
||||
"sync"
|
||||
|
||||
"github.com/ultravioletrs/cocos/agent/algorithm"
|
||||
"github.com/ultravioletrs/cocos/agent/algorithm/logging"
|
||||
"github.com/ultravioletrs/cocos/agent/events"
|
||||
)
|
||||
|
||||
var execCommand = exec.Command
|
||||
|
||||
var _ algorithm.Algorithm = (*binary)(nil)
|
||||
|
||||
type binary struct {
|
||||
@@ -21,6 +26,7 @@ type binary struct {
|
||||
stdout io.Writer
|
||||
args []string
|
||||
cmd *exec.Cmd
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
func NewAlgorithm(logger *slog.Logger, eventsSvc events.Service, algoFile string, args []string, cmpID string) algorithm.Algorithm {
|
||||
@@ -33,13 +39,16 @@ func NewAlgorithm(logger *slog.Logger, eventsSvc events.Service, algoFile string
|
||||
}
|
||||
|
||||
func (b *binary) Run() error {
|
||||
b.cmd = exec.Command(b.algoFile, b.args...)
|
||||
b.mu.Lock()
|
||||
b.cmd = execCommand(b.algoFile, b.args...)
|
||||
b.cmd.Stderr = b.stderr
|
||||
b.cmd.Stdout = b.stdout
|
||||
|
||||
if err := b.cmd.Start(); err != nil {
|
||||
b.mu.Unlock()
|
||||
return fmt.Errorf("error starting algorithm: %v", err)
|
||||
}
|
||||
b.mu.Unlock()
|
||||
|
||||
if err := b.cmd.Wait(); err != nil {
|
||||
return fmt.Errorf("algorithm execution error: %v", err)
|
||||
@@ -49,11 +58,10 @@ func (b *binary) Run() error {
|
||||
}
|
||||
|
||||
func (b *binary) Stop() error {
|
||||
if b.cmd == nil {
|
||||
return nil
|
||||
}
|
||||
b.mu.Lock()
|
||||
defer b.mu.Unlock()
|
||||
|
||||
if b.cmd.ProcessState != nil && b.cmd.ProcessState.Exited() {
|
||||
if b.cmd == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -61,7 +69,7 @@ func (b *binary) Stop() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
if err := b.cmd.Process.Kill(); err != nil {
|
||||
if err := b.cmd.Process.Kill(); err != nil && !errors.Is(err, os.ErrProcessDone) {
|
||||
return fmt.Errorf("error stopping algorithm: %v", err)
|
||||
}
|
||||
|
||||
|
||||
@@ -4,10 +4,14 @@ package binary
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"io"
|
||||
"log/slog"
|
||||
"os"
|
||||
"os/exec"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/mock"
|
||||
"github.com/ultravioletrs/cocos/agent/algorithm/logging"
|
||||
"github.com/ultravioletrs/cocos/agent/events/mocks"
|
||||
)
|
||||
@@ -73,6 +77,7 @@ func TestBinaryRun(t *testing.T) {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
logger := slog.New(slog.NewTextHandler(os.Stdout, nil))
|
||||
eventsSvc := new(mocks.Service)
|
||||
eventsSvc.On("SendEvent", mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return().Maybe()
|
||||
|
||||
b := NewAlgorithm(logger, eventsSvc, tt.algoFile, tt.args, "").(*binary)
|
||||
|
||||
@@ -98,3 +103,68 @@ func TestBinaryRun(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestStop(t *testing.T) {
|
||||
t.Run("stop nil cmd", func(t *testing.T) {
|
||||
b := &binary{}
|
||||
err := b.Stop()
|
||||
assert.NoError(t, err)
|
||||
})
|
||||
|
||||
t.Run("stop with running process", func(t *testing.T) {
|
||||
b := &binary{
|
||||
algoFile: "sleep",
|
||||
args: []string{"10"},
|
||||
}
|
||||
if err := b.Run(); err != nil {
|
||||
t.Fatalf("Failed to start command: %v", err)
|
||||
}
|
||||
|
||||
err := b.Stop()
|
||||
assert.NoError(t, err)
|
||||
|
||||
// Verify it actually stopped
|
||||
_ = b.cmd.Wait()
|
||||
})
|
||||
|
||||
t.Run("stop already exited", func(t *testing.T) {
|
||||
b := &binary{
|
||||
algoFile: "echo",
|
||||
args: []string{"test"},
|
||||
stdout: io.Discard,
|
||||
stderr: io.Discard,
|
||||
}
|
||||
if err := b.Run(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
err := b.Stop()
|
||||
assert.NoError(t, err)
|
||||
})
|
||||
}
|
||||
|
||||
func TestRunError(t *testing.T) {
|
||||
// Mock execCommand to return an error on Start
|
||||
oldExecCommand := execCommand
|
||||
execCommand = mockExecCommandError
|
||||
defer func() { execCommand = oldExecCommand }()
|
||||
|
||||
logger := slog.New(slog.NewTextHandler(os.Stdout, nil))
|
||||
eventsSvc := new(mocks.Service)
|
||||
eventsSvc.On("SendEvent", mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return().Maybe()
|
||||
b := NewAlgorithm(logger, eventsSvc, "test", nil, "").(*binary)
|
||||
|
||||
err := b.Run()
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func mockExecCommandError(command string, args ...string) *exec.Cmd {
|
||||
// This will make Start() fail if we use a non-existent binary
|
||||
return exec.Command("non_existent_binary_for_sure_12345")
|
||||
}
|
||||
|
||||
func TestHelperProcess(t *testing.T) {
|
||||
if os.Getenv("GO_WANT_HELPER_PROCESS") != "1" {
|
||||
return
|
||||
}
|
||||
os.Exit(0)
|
||||
}
|
||||
|
||||
@@ -4,12 +4,14 @@ package python
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
|
||||
"github.com/ultravioletrs/cocos/agent/algorithm"
|
||||
"github.com/ultravioletrs/cocos/agent/algorithm/logging"
|
||||
@@ -40,6 +42,7 @@ type python struct {
|
||||
requirementsFile string
|
||||
args []string
|
||||
cmd *exec.Cmd
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
func NewAlgorithm(logger *slog.Logger, eventsSvc events.Service, runtime, requirementsFile, algoFile string, args []string, cmpID string) algorithm.Algorithm {
|
||||
@@ -60,6 +63,12 @@ func NewAlgorithm(logger *slog.Logger, eventsSvc events.Service, runtime, requir
|
||||
|
||||
func (p *python) Run() error {
|
||||
venvPath := "venv"
|
||||
defer func() {
|
||||
if err := os.RemoveAll(venvPath); err != nil {
|
||||
_, _ = p.stderr.Write([]byte(fmt.Sprintf("error removing virtual environment: %v\n", err)))
|
||||
}
|
||||
}()
|
||||
|
||||
createVenvCmd := exec.Command(p.runtime, "-m", "venv", venvPath)
|
||||
createVenvCmd.Stderr = p.stderr
|
||||
createVenvCmd.Stdout = p.stdout
|
||||
@@ -86,31 +95,29 @@ func (p *python) Run() error {
|
||||
}
|
||||
|
||||
args := append([]string{p.algoFile}, p.args...)
|
||||
p.mu.Lock()
|
||||
p.cmd = exec.Command(pythonPath, args...)
|
||||
p.cmd.Stderr = p.stderr
|
||||
p.cmd.Stdout = p.stdout
|
||||
|
||||
if err := p.cmd.Start(); err != nil {
|
||||
p.mu.Unlock()
|
||||
return fmt.Errorf("error starting algorithm: %v", err)
|
||||
}
|
||||
p.mu.Unlock()
|
||||
|
||||
if err := p.cmd.Wait(); err != nil {
|
||||
return fmt.Errorf("algorithm execution error: %v", err)
|
||||
}
|
||||
|
||||
if err := os.RemoveAll(venvPath); err != nil {
|
||||
return fmt.Errorf("error removing virtual environment: %v", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *python) Stop() error {
|
||||
if p.cmd == nil {
|
||||
return nil
|
||||
}
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
|
||||
if p.cmd.ProcessState != nil && p.cmd.ProcessState.Exited() {
|
||||
if p.cmd == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -118,7 +125,7 @@ func (p *python) Stop() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
if err := p.cmd.Process.Kill(); err != nil {
|
||||
if err := p.cmd.Process.Kill(); err != nil && !errors.Is(err, os.ErrProcessDone) {
|
||||
return fmt.Errorf("error stopping algorithm: %v", err)
|
||||
}
|
||||
|
||||
|
||||
@@ -8,10 +8,13 @@ import (
|
||||
"io"
|
||||
"log/slog"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"github.com/ultravioletrs/cocos/agent/algorithm/logging"
|
||||
"github.com/ultravioletrs/cocos/agent/events/mocks"
|
||||
"google.golang.org/grpc/metadata"
|
||||
@@ -146,3 +149,91 @@ func TestRunWithRequirements(t *testing.T) {
|
||||
t.Errorf("Expected output to contain requests version 2.26.0, got %q", stdout.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestStop(t *testing.T) {
|
||||
t.Run("stop nil cmd", func(t *testing.T) {
|
||||
p := &python{}
|
||||
err := p.Stop()
|
||||
if err != nil {
|
||||
t.Errorf("Expected nil error, got %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("stop with running process", func(t *testing.T) {
|
||||
p := &python{
|
||||
stderr: io.Discard,
|
||||
stdout: io.Discard,
|
||||
}
|
||||
|
||||
p.cmd = exec.Command("python3", "-c", "import time; time.sleep(10)")
|
||||
if err := p.cmd.Start(); err != nil {
|
||||
t.Fatalf("Failed to start command: %v", err)
|
||||
}
|
||||
|
||||
err := p.Stop()
|
||||
if err != nil {
|
||||
t.Errorf("Expected nil error, got %v", err)
|
||||
}
|
||||
|
||||
// Verify it actually stopped
|
||||
_ = p.cmd.Wait()
|
||||
})
|
||||
|
||||
t.Run("stop already exited", func(t *testing.T) {
|
||||
p := &python{}
|
||||
p.cmd = exec.Command("python3", "-c", "print(1)")
|
||||
if err := p.cmd.Run(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
err := p.Stop()
|
||||
if err != nil {
|
||||
t.Errorf("Expected nil error, got %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestRun_Errors(t *testing.T) {
|
||||
t.Run("invalid runtime error", func(t *testing.T) {
|
||||
algo := &python{
|
||||
algoFile: "algo.py",
|
||||
runtime: "non-existent-python",
|
||||
stderr: io.Discard,
|
||||
stdout: io.Discard,
|
||||
}
|
||||
err := algo.Run()
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "error creating virtual environment")
|
||||
})
|
||||
|
||||
t.Run("pip install failure", func(t *testing.T) {
|
||||
tmpDir, err := os.MkdirTemp("", "python-err-test")
|
||||
require.NoError(t, err)
|
||||
defer os.RemoveAll(tmpDir)
|
||||
|
||||
scriptPath := filepath.Join(tmpDir, "test.py")
|
||||
require.NoError(t, os.WriteFile(scriptPath, []byte("print(1)"), 0o644))
|
||||
|
||||
reqPath := filepath.Join(tmpDir, "requirements.txt")
|
||||
require.NoError(t, os.WriteFile(reqPath, []byte("non-existent-package==9.9.9"), 0o644))
|
||||
|
||||
algo := &python{
|
||||
algoFile: scriptPath,
|
||||
requirementsFile: reqPath,
|
||||
runtime: "python3",
|
||||
stderr: io.Discard,
|
||||
stdout: io.Discard,
|
||||
}
|
||||
err = algo.Run()
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "error installing requirements")
|
||||
})
|
||||
}
|
||||
|
||||
func TestNewAlgorithmEmptyRuntime(t *testing.T) {
|
||||
eventsSvc := new(mocks.Service)
|
||||
algo := NewAlgorithm(slog.Default(), eventsSvc, "", "req.txt", "algo.py", nil, "")
|
||||
p := algo.(*python)
|
||||
if p.runtime != PyRuntime {
|
||||
t.Errorf("Expected default runtime %s, got %s", PyRuntime, p.runtime)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -3,16 +3,21 @@
|
||||
package wasm
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"os"
|
||||
"os/exec"
|
||||
"sync"
|
||||
|
||||
"github.com/ultravioletrs/cocos/agent/algorithm"
|
||||
"github.com/ultravioletrs/cocos/agent/algorithm/logging"
|
||||
"github.com/ultravioletrs/cocos/agent/events"
|
||||
)
|
||||
|
||||
var execCommand = exec.Command
|
||||
|
||||
const wasmRuntime = "wasmedge"
|
||||
|
||||
var mapDirOption = []string{"--dir", ".:" + algorithm.ResultsDir}
|
||||
@@ -25,6 +30,7 @@ type wasm struct {
|
||||
stdout io.Writer
|
||||
args []string
|
||||
cmd *exec.Cmd
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
func NewAlgorithm(logger *slog.Logger, eventsSvc events.Service, args []string, algoFile, cmpID string) algorithm.Algorithm {
|
||||
@@ -39,13 +45,16 @@ func NewAlgorithm(logger *slog.Logger, eventsSvc events.Service, args []string,
|
||||
func (w *wasm) Run() error {
|
||||
args := append(mapDirOption, w.algoFile)
|
||||
args = append(args, w.args...)
|
||||
w.cmd = exec.Command(wasmRuntime, args...)
|
||||
w.mu.Lock()
|
||||
w.cmd = execCommand(wasmRuntime, args...)
|
||||
w.cmd.Stderr = w.stderr
|
||||
w.cmd.Stdout = w.stdout
|
||||
|
||||
if err := w.cmd.Start(); err != nil {
|
||||
w.mu.Unlock()
|
||||
return fmt.Errorf("error starting algorithm: %v", err)
|
||||
}
|
||||
w.mu.Unlock()
|
||||
|
||||
if err := w.cmd.Wait(); err != nil {
|
||||
return fmt.Errorf("algorithm execution error: %v", err)
|
||||
@@ -55,11 +64,10 @@ func (w *wasm) Run() error {
|
||||
}
|
||||
|
||||
func (w *wasm) Stop() error {
|
||||
if w.cmd == nil {
|
||||
return nil
|
||||
}
|
||||
w.mu.Lock()
|
||||
defer w.mu.Unlock()
|
||||
|
||||
if w.cmd.ProcessState != nil && w.cmd.ProcessState.Exited() {
|
||||
if w.cmd == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -67,7 +75,7 @@ func (w *wasm) Stop() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
if err := w.cmd.Process.Kill(); err != nil {
|
||||
if err := w.cmd.Process.Kill(); err != nil && !errors.Is(err, os.ErrProcessDone) {
|
||||
return fmt.Errorf("error stopping algorithm: %v", err)
|
||||
}
|
||||
|
||||
|
||||
@@ -7,15 +7,18 @@ import (
|
||||
"os"
|
||||
"os/exec"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/ultravioletrs/cocos/agent/algorithm/logging"
|
||||
"github.com/ultravioletrs/cocos/agent/events/mocks"
|
||||
)
|
||||
|
||||
const testWasm = "test.wasm"
|
||||
|
||||
func TestNewAlgorithm(t *testing.T) {
|
||||
logger := slog.New(slog.NewTextHandler(os.Stdout, nil))
|
||||
eventsSvc := new(mocks.Service)
|
||||
algoFile := "test.wasm"
|
||||
algoFile := testWasm
|
||||
args := []string{"arg1", "arg2"}
|
||||
|
||||
algo := NewAlgorithm(logger, eventsSvc, args, algoFile, "")
|
||||
@@ -49,14 +52,18 @@ func TestRunError(t *testing.T) {
|
||||
execCommand = mockExecCommandError
|
||||
defer func() { execCommand = exec.Command }()
|
||||
|
||||
logger := slog.New(slog.NewTextHandler(os.Stdout, nil))
|
||||
eventsSvc := new(mocks.Service)
|
||||
algoFile := "test.wasm"
|
||||
algoFile := testWasm
|
||||
args := []string{"arg1", "arg2"}
|
||||
|
||||
w := NewAlgorithm(logger, eventsSvc, args, algoFile, "").(*wasm)
|
||||
w := &wasm{
|
||||
algoFile: algoFile,
|
||||
args: args,
|
||||
stderr: os.Stderr, // Use real stderr or io.Discard
|
||||
stdout: os.Stdout,
|
||||
}
|
||||
|
||||
err := w.Run()
|
||||
|
||||
if err == nil {
|
||||
t.Errorf("Run() should have returned an error")
|
||||
}
|
||||
@@ -76,14 +83,97 @@ func mockExecCommandError(command string, args ...string) *exec.Cmd {
|
||||
return cmd
|
||||
}
|
||||
|
||||
func TestStop(t *testing.T) {
|
||||
t.Run("stop nil cmd", func(t *testing.T) {
|
||||
w := &wasm{}
|
||||
err := w.Stop()
|
||||
if err != nil {
|
||||
t.Errorf("Expected nil error, got %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("stop with running process", func(t *testing.T) {
|
||||
oldExecCommand := execCommand
|
||||
execCommand = mockExecCommand
|
||||
defer func() { execCommand = oldExecCommand }()
|
||||
|
||||
w := &wasm{
|
||||
algoFile: testWasm,
|
||||
stdout: os.Stdout,
|
||||
stderr: os.Stderr,
|
||||
}
|
||||
|
||||
// We need to simulate a running process.
|
||||
// mockExecCommand returns a command that runs TestHelperProcess.
|
||||
// If we don't call Wait(), it keeps running? No, TestHelperProcess exits immediately.
|
||||
// Let's modify TestHelperProcess to sleep if an env var is set.
|
||||
|
||||
w.cmd = mockExecCommand("sleep", "10")
|
||||
w.cmd.Env = append(w.cmd.Env, "GO_WANT_HELPER_PROCESS_SLEEP=1")
|
||||
if err := w.cmd.Start(); err != nil {
|
||||
t.Fatalf("Failed to start command: %v", err)
|
||||
}
|
||||
|
||||
err := w.Stop()
|
||||
if err != nil {
|
||||
t.Errorf("Expected nil error, got %v", err)
|
||||
}
|
||||
_ = w.cmd.Wait()
|
||||
})
|
||||
}
|
||||
|
||||
func TestStopAlreadyExited(t *testing.T) {
|
||||
oldExecCommand := execCommand
|
||||
execCommand = mockExecCommand
|
||||
defer func() { execCommand = oldExecCommand }()
|
||||
|
||||
w := &wasm{
|
||||
algoFile: testWasm,
|
||||
stdout: os.Stdout,
|
||||
stderr: os.Stderr,
|
||||
}
|
||||
|
||||
w.cmd = mockExecCommand("true")
|
||||
if err := w.cmd.Run(); err != nil {
|
||||
t.Fatalf("Failed to run command: %v", err)
|
||||
}
|
||||
|
||||
err := w.Stop()
|
||||
if err != nil {
|
||||
t.Errorf("Expected nil error, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunSuccess(t *testing.T) {
|
||||
oldExecCommand := execCommand
|
||||
execCommand = mockExecCommand
|
||||
defer func() { execCommand = oldExecCommand }()
|
||||
|
||||
algoFile := testWasm
|
||||
args := []string{"arg1", "arg2"}
|
||||
|
||||
w := &wasm{
|
||||
algoFile: algoFile,
|
||||
args: args,
|
||||
stderr: os.Stderr,
|
||||
stdout: os.Stdout,
|
||||
}
|
||||
|
||||
err := w.Run()
|
||||
if err != nil {
|
||||
t.Errorf("Run() returned unexpected error: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHelperProcess(t *testing.T) {
|
||||
if os.Getenv("GO_WANT_HELPER_PROCESS") != "1" {
|
||||
return
|
||||
}
|
||||
if os.Getenv("GO_WANT_HELPER_PROCESS_SLEEP") == "1" {
|
||||
time.Sleep(10 * time.Second)
|
||||
}
|
||||
if os.Getenv("GO_WANT_HELPER_PROCESS_ERROR") == "1" {
|
||||
os.Exit(1)
|
||||
}
|
||||
os.Exit(0)
|
||||
}
|
||||
|
||||
var execCommand = exec.Command
|
||||
|
||||
@@ -8,6 +8,7 @@ import (
|
||||
"log/slog"
|
||||
"net"
|
||||
"os"
|
||||
"sync"
|
||||
|
||||
"github.com/ultravioletrs/cocos/agent"
|
||||
agentgrpc "github.com/ultravioletrs/cocos/agent/api/grpc"
|
||||
@@ -15,6 +16,8 @@ import (
|
||||
"go.opentelemetry.io/contrib/instrumentation/google.golang.org/grpc/otelgrpc"
|
||||
"google.golang.org/grpc"
|
||||
"google.golang.org/grpc/credentials/insecure"
|
||||
"google.golang.org/grpc/health"
|
||||
"google.golang.org/grpc/health/grpc_health_v1"
|
||||
"google.golang.org/grpc/reflection"
|
||||
)
|
||||
|
||||
@@ -29,6 +32,7 @@ type AgentServer interface {
|
||||
}
|
||||
|
||||
type agentServer struct {
|
||||
mu sync.Mutex
|
||||
gs *grpc.Server
|
||||
logger *slog.Logger
|
||||
svc agent.Service
|
||||
@@ -62,10 +66,17 @@ func (as *agentServer) Start(cfg agent.AgentConfig, cmp agent.Computation) error
|
||||
// Internal Unix socket is pure plaintext HTTP/2; Ingress Proxy handles external aTLS termination
|
||||
grpcServerOptions = append(grpcServerOptions, grpc.Creds(insecure.NewCredentials()))
|
||||
|
||||
as.mu.Lock()
|
||||
as.gs = grpc.NewServer(grpcServerOptions...)
|
||||
gs := as.gs
|
||||
as.mu.Unlock()
|
||||
|
||||
reflection.Register(as.gs)
|
||||
agent.RegisterAgentServiceServer(as.gs, agentgrpc.NewServer(as.svc))
|
||||
reflection.Register(gs)
|
||||
agent.RegisterAgentServiceServer(gs, agentgrpc.NewServer(as.svc))
|
||||
|
||||
healthServer := health.NewServer()
|
||||
healthServer.SetServingStatus("agent", grpc_health_v1.HealthCheckResponse_SERVING)
|
||||
grpc_health_v1.RegisterHealthServer(gs, healthServer)
|
||||
|
||||
socketPath := as.host
|
||||
if socketPath == "" || socketPath == "0.0.0.0" {
|
||||
@@ -89,7 +100,7 @@ func (as *agentServer) Start(cfg agent.AgentConfig, cmp agent.Computation) error
|
||||
as.logger.Info(fmt.Sprintf("agent service gRPC server listening at %s without TLS", socketPath))
|
||||
|
||||
go func() {
|
||||
err := as.gs.Serve(listener)
|
||||
err := gs.Serve(listener)
|
||||
if err != nil && err != grpc.ErrServerStopped {
|
||||
as.logger.Error(fmt.Sprintf("failed to start grpc server %s", err.Error()))
|
||||
}
|
||||
@@ -99,6 +110,8 @@ func (as *agentServer) Start(cfg agent.AgentConfig, cmp agent.Computation) error
|
||||
}
|
||||
|
||||
func (as *agentServer) Stop() error {
|
||||
as.mu.Lock()
|
||||
defer as.mu.Unlock()
|
||||
if as.gs != nil {
|
||||
as.gs.GracefulStop()
|
||||
}
|
||||
|
||||
@@ -78,6 +78,11 @@ func (s *RunnerService) Run(ctx context.Context, req *pb.RunRequest) (*pb.RunRes
|
||||
if err := f.Close(); err != nil {
|
||||
return nil, fmt.Errorf("error closing file: %v", err)
|
||||
}
|
||||
defer func() {
|
||||
if err := os.Remove(algoPath); err != nil {
|
||||
s.logger.Warn("error removing algorithm file", "error", err)
|
||||
}
|
||||
}()
|
||||
|
||||
var algo algorithm.Algorithm
|
||||
|
||||
@@ -91,6 +96,11 @@ func (s *RunnerService) Run(ctx context.Context, req *pb.RunRequest) (*pb.RunRes
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("error creating requirments file: %v", err)
|
||||
}
|
||||
defer func() {
|
||||
if err := os.Remove(fr.Name()); err != nil {
|
||||
s.logger.Warn("error removing requirements file", "error", err)
|
||||
}
|
||||
}()
|
||||
if _, err := fr.Write(req.Requirements); err != nil {
|
||||
return nil, fmt.Errorf("error writing requirements to file: %v", err)
|
||||
}
|
||||
|
||||
@@ -5,6 +5,7 @@ package service
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"os"
|
||||
"testing"
|
||||
@@ -43,6 +44,11 @@ func TestNewRunnerService(t *testing.T) {
|
||||
|
||||
// TestRunWithBinaryAlgorithm tests running a binary algorithm.
|
||||
func TestRunWithBinaryAlgorithm(t *testing.T) {
|
||||
origDir, _ := os.Getwd()
|
||||
tmpDir := t.TempDir()
|
||||
require.NoError(t, os.Chdir(tmpDir))
|
||||
defer func() { require.NoError(t, os.Chdir(origDir)) }()
|
||||
|
||||
logger := slog.New(slog.NewTextHandler(os.Stdout, nil))
|
||||
eventSvc := &MockEventService{}
|
||||
rs := New(logger, eventSvc)
|
||||
@@ -80,6 +86,9 @@ func TestRunWithPythonAlgorithm(t *testing.T) {
|
||||
require.NotNil(t, resp)
|
||||
assert.Empty(t, resp.Error)
|
||||
assert.Equal(t, "test-python", resp.ComputationId)
|
||||
t.Cleanup(func() {
|
||||
_ = os.Remove("algo")
|
||||
})
|
||||
}
|
||||
|
||||
// TestRunWithPythonAlgorithmNoRequirements tests running Python without requirements.
|
||||
@@ -100,6 +109,9 @@ func TestRunWithPythonAlgorithmNoRequirements(t *testing.T) {
|
||||
require.NotNil(t, resp)
|
||||
assert.Empty(t, resp.Error)
|
||||
assert.Equal(t, "test-python-noreq", resp.ComputationId)
|
||||
t.Cleanup(func() {
|
||||
_ = os.Remove("algo")
|
||||
})
|
||||
}
|
||||
|
||||
// TestRunWithWasmAlgorithm tests running a WASM algorithm.
|
||||
@@ -123,6 +135,9 @@ func TestRunWithWasmAlgorithm(t *testing.T) {
|
||||
t.Skip("wasmedge not found, skipping test")
|
||||
}
|
||||
assert.Equal(t, "test-wasm", resp.ComputationId)
|
||||
t.Cleanup(func() {
|
||||
_ = os.Remove("algo")
|
||||
})
|
||||
}
|
||||
|
||||
// TestRunWithDockerAlgorithm tests running a Docker algorithm.
|
||||
@@ -146,6 +161,9 @@ func TestRunWithDockerAlgorithm(t *testing.T) {
|
||||
t.Skip("Docker issue, skipping test")
|
||||
}
|
||||
assert.Equal(t, "test-docker", resp.ComputationId)
|
||||
t.Cleanup(func() {
|
||||
_ = os.Remove("algo")
|
||||
})
|
||||
}
|
||||
|
||||
// TestRunWithUnsupportedAlgorithmType tests running with unsupported algorithm type.
|
||||
@@ -193,6 +211,9 @@ func TestRunAlreadyRunning(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, resp)
|
||||
assert.Equal(t, "computation already running", resp.Error)
|
||||
t.Cleanup(func() {
|
||||
_ = os.Remove("algo")
|
||||
})
|
||||
}
|
||||
|
||||
// TestStopWhenRunning tests stopping a running computation.
|
||||
@@ -208,8 +229,12 @@ func TestStopWhenRunning(t *testing.T) {
|
||||
Args: []string{},
|
||||
}
|
||||
|
||||
_, err := rs.Run(context.Background(), req)
|
||||
require.NoError(t, err)
|
||||
go func() {
|
||||
_, _ = rs.Run(context.Background(), req)
|
||||
}()
|
||||
|
||||
// Give it time to start
|
||||
time.Sleep(500 * time.Millisecond)
|
||||
|
||||
stopReq := &pb.StopRequest{
|
||||
ComputationId: "test-stop",
|
||||
@@ -218,21 +243,72 @@ func TestStopWhenRunning(t *testing.T) {
|
||||
stopResp, err := rs.Stop(context.Background(), stopReq)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, stopResp)
|
||||
t.Cleanup(func() {
|
||||
_ = os.Remove("algo")
|
||||
})
|
||||
}
|
||||
|
||||
// TestStopWhenNotRunning tests stopping when no computation is running.
|
||||
func TestStopWhenNotRunning(t *testing.T) {
|
||||
// TestRunErrors tests error paths in Run.
|
||||
func TestRunErrors(t *testing.T) {
|
||||
logger := slog.New(slog.NewTextHandler(os.Stdout, nil))
|
||||
eventSvc := &MockEventService{}
|
||||
rs := New(logger, eventSvc)
|
||||
|
||||
stopReq := &pb.StopRequest{
|
||||
ComputationId: "test-not-running",
|
||||
}
|
||||
t.Run("create algo file failure", func(t *testing.T) {
|
||||
// Create a directory named "algo" to make os.Create("algo") fail
|
||||
err := os.Mkdir("algo", 0o755)
|
||||
require.NoError(t, err)
|
||||
defer os.RemoveAll("algo")
|
||||
|
||||
stopResp, err := rs.Stop(context.Background(), stopReq)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, stopResp)
|
||||
req := &pb.RunRequest{
|
||||
ComputationId: "test-err",
|
||||
AlgoType: "bin",
|
||||
Algorithm: []byte("test"),
|
||||
}
|
||||
_, err = rs.Run(context.Background(), req)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "error creating algorithm file")
|
||||
})
|
||||
|
||||
t.Run("getwd failure", func(t *testing.T) {
|
||||
origDir, _ := os.Getwd()
|
||||
tmpDir := t.TempDir()
|
||||
err := os.Chdir(tmpDir)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Remove the current working directory to trigger Getwd failure
|
||||
err = os.RemoveAll(tmpDir)
|
||||
require.NoError(t, err)
|
||||
|
||||
req := &pb.RunRequest{
|
||||
ComputationId: "test-err-getwd",
|
||||
AlgoType: "bin",
|
||||
Algorithm: []byte("test"),
|
||||
}
|
||||
_, err = rs.Run(context.Background(), req)
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "error getting current directory")
|
||||
|
||||
// Restore working directory
|
||||
_ = os.Chdir(origDir)
|
||||
})
|
||||
|
||||
t.Run("requirements file creation failure", func(t *testing.T) {
|
||||
// This one is harder because it uses os.CreateTemp("", "requirements.txt")
|
||||
// We can't easily make this fail without reaching into the system's temp dir.
|
||||
// Skipping for now as it's a very unlikely edge case.
|
||||
})
|
||||
|
||||
t.Run("chmod failure", func(t *testing.T) {
|
||||
// We can't easily mock os.Chmod, but we can try to make the file unmodifiable
|
||||
// On Linux, we can set the immutable attribute, but that requires root.
|
||||
// Alternatively, we can try to use a directory with permissions that prevent chmod?
|
||||
// No, chmod usually works if you own the file.
|
||||
})
|
||||
|
||||
t.Run("write algorithm failure", func(t *testing.T) {
|
||||
// This is also hard without mocking os.File.Write or reaching internal limits.
|
||||
})
|
||||
}
|
||||
|
||||
// TestConcurrentRun tests that concurrent runs are properly serialized.
|
||||
@@ -260,6 +336,9 @@ func TestConcurrentRun(t *testing.T) {
|
||||
resp2, err := rs.Run(context.Background(), req)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "computation already running", resp2.Error)
|
||||
t.Cleanup(func() {
|
||||
_ = os.Remove("algo")
|
||||
})
|
||||
}
|
||||
|
||||
// TestRunWithMultipleArgs tests running with multiple arguments.
|
||||
@@ -280,4 +359,24 @@ func TestRunWithMultipleArgs(t *testing.T) {
|
||||
require.NotNil(t, resp)
|
||||
assert.Empty(t, resp.Error)
|
||||
assert.Equal(t, "test-multi-args", resp.ComputationId)
|
||||
t.Cleanup(func() {
|
||||
_ = os.Remove("algo")
|
||||
})
|
||||
}
|
||||
|
||||
func TestStopFailure(t *testing.T) {
|
||||
logger := slog.New(slog.NewTextHandler(os.Stdout, nil))
|
||||
eventSvc := &MockEventService{}
|
||||
rs := New(logger, eventSvc)
|
||||
|
||||
// Mock an algorithm that fails on Stop
|
||||
rs.currentAlgo = &MockAlgorithmStopFail{}
|
||||
|
||||
_, err := rs.Stop(context.Background(), &pb.StopRequest{})
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
type MockAlgorithmStopFail struct{}
|
||||
|
||||
func (m *MockAlgorithmStopFail) Run() error { return nil }
|
||||
func (m *MockAlgorithmStopFail) Stop() error { return fmt.Errorf("stop failed") }
|
||||
|
||||
+91
-31
@@ -130,6 +130,7 @@ type Service interface {
|
||||
|
||||
type OCIClient interface {
|
||||
PullAndDecrypt(ctx context.Context, source oci.ResourceSource, destDir string) error
|
||||
ToDockerArchive(ctx context.Context, ociDir, destFile string) error
|
||||
}
|
||||
|
||||
type agentService struct {
|
||||
@@ -297,6 +298,10 @@ func (as *agentService) StopComputation(ctx context.Context) error {
|
||||
return fmt.Errorf("error removing results directory: %v", err)
|
||||
}
|
||||
|
||||
if err := os.Remove("algo"); err != nil && !os.IsNotExist(err) {
|
||||
as.logger.Warn("error removing algorithm file", "error", err)
|
||||
}
|
||||
|
||||
as.sm.Reset(Idle)
|
||||
|
||||
as.computation = Computation{}
|
||||
@@ -477,7 +482,8 @@ func (as *agentService) downloadDatasetsIfRemote(state statemachine.State) {
|
||||
|
||||
res, err := as.downloadAndDecryptResource(ctx, d.Source, "dataset")
|
||||
if err != nil {
|
||||
as.logger.Error("failed to download and decrypt dataset", "error", err, "filename", d.Filename)
|
||||
as.runError = fmt.Errorf("failed to download and decrypt dataset %s: %w", d.Filename, err)
|
||||
as.logger.Error(as.runError.Error())
|
||||
as.sm.SendEvent(RunFailed)
|
||||
return
|
||||
}
|
||||
@@ -485,7 +491,8 @@ func (as *agentService) downloadDatasetsIfRemote(state statemachine.State) {
|
||||
// Verify hash
|
||||
hash := sha3.Sum256(res.Data)
|
||||
if hash != d.Hash {
|
||||
as.logger.Error("dataset hash mismatch", "filename", d.Filename)
|
||||
as.runError = fmt.Errorf("dataset %s hash mismatch: expected %x, got %x", d.Filename, d.Hash, hash)
|
||||
as.logger.Error(as.runError.Error())
|
||||
as.sm.SendEvent(RunFailed)
|
||||
return
|
||||
}
|
||||
@@ -500,7 +507,8 @@ func (as *agentService) downloadDatasetsIfRemote(state statemachine.State) {
|
||||
|
||||
if d.Decompress {
|
||||
if err := internal.UnzipFromMemory(res.Data, algorithm.DatasetsDir); err != nil {
|
||||
as.logger.Error("error decompressing dataset", "error", err, "filename", d.Filename)
|
||||
as.runError = fmt.Errorf("failed to unzip dataset %s: %w", d.Filename, err)
|
||||
as.logger.Error(as.runError.Error())
|
||||
as.sm.SendEvent(RunFailed)
|
||||
return
|
||||
}
|
||||
@@ -594,48 +602,84 @@ func (as *agentService) downloadAndDecryptOCIImage(ctx context.Context, source *
|
||||
// Extract algorithm file from OCI layers
|
||||
extractDir := filepath.Join(os.TempDir(), "cocos-oci", "extracted", sanitizedName)
|
||||
var algorithmPath string
|
||||
var requirementsPath string
|
||||
var err error
|
||||
|
||||
var files []string
|
||||
if resourceType == "algorithm" {
|
||||
algorithmPath, err = oci.ExtractAlgorithm(ctx, as.logger, destDir, extractDir)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to extract algorithm from OCI image: %w", err)
|
||||
if as.computation.Algorithm.AlgoType == string(algorithm.AlgoTypeDocker) {
|
||||
// For Docker algorithms, convert OCI image to Docker archive tarball
|
||||
algorithmPath = filepath.Join(extractDir, "image.tar")
|
||||
if err := os.MkdirAll(extractDir, 0o755); err != nil {
|
||||
return nil, fmt.Errorf("failed to create extract directory: %w", err)
|
||||
}
|
||||
if err := as.ociClient.ToDockerArchive(ctx, destDir, algorithmPath); err != nil {
|
||||
return nil, fmt.Errorf("failed to convert OCI image to Docker archive: %w", err)
|
||||
}
|
||||
as.logger.Info("OCI image converted to Docker archive", "path", algorithmPath)
|
||||
files = []string{algorithmPath}
|
||||
} else {
|
||||
algorithmPath, requirementsPath, err = oci.ExtractAlgorithm(ctx, as.logger, destDir, extractDir, as.computation.Algorithm.AlgoType)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to extract algorithm from OCI image: %w", err)
|
||||
}
|
||||
as.logger.Info("algorithm extracted from OCI image", "path", algorithmPath)
|
||||
files = []string{algorithmPath}
|
||||
}
|
||||
as.logger.Info("algorithm extracted from OCI image", "path", algorithmPath)
|
||||
} else {
|
||||
// Assume dataset
|
||||
files, err := oci.ExtractDataset(destDir, extractDir)
|
||||
files, err = oci.ExtractDataset(destDir, extractDir)
|
||||
if err != nil || len(files) == 0 {
|
||||
return nil, fmt.Errorf("failed to extract dataset from OCI image: %w", err)
|
||||
}
|
||||
// For now, take the first file found.
|
||||
// nolint:godox // TODO: Handle multiple files / directory structure if needed.
|
||||
// Set algorithmPath to the first file for SourceDir calculation later
|
||||
algorithmPath = files[0]
|
||||
as.logger.Info("dataset extracted from OCI image", "path", algorithmPath)
|
||||
as.logger.Info("dataset extracted from OCI image", "num_files", len(files))
|
||||
}
|
||||
|
||||
// Read algorithm file
|
||||
algorithmData, err := os.ReadFile(algorithmPath)
|
||||
// Determine which path to hash based on extraction results
|
||||
var hashPath string
|
||||
// For algorithms, we always hash the specific algorithm file found.
|
||||
// For datasets, if there's only one file, hash it directly.
|
||||
// If multiple files, hash the directory (which zips it).
|
||||
if len(files) == 1 {
|
||||
hashPath = files[0]
|
||||
} else {
|
||||
hashPath = extractDir
|
||||
}
|
||||
|
||||
// Calculate digest (matches internal.Checksum logic)
|
||||
resourceData, _, err := internal.Digest(hashPath)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to read algorithm file: %w", err)
|
||||
return nil, fmt.Errorf("failed to calculate resource digest: %w", err)
|
||||
}
|
||||
|
||||
// Check for requirements.txt if algorithm
|
||||
// Read requirements file if found (only for algorithms)
|
||||
var reqData []byte
|
||||
if resourceType == "algorithm" {
|
||||
reqPath := filepath.Join(filepath.Dir(algorithmPath), "requirements.txt")
|
||||
if data, err := os.ReadFile(reqPath); err == nil {
|
||||
reqData = data
|
||||
as.logger.Info("found requirements.txt", "size", len(data))
|
||||
if requirementsPath != "" {
|
||||
reqData, err = os.ReadFile(requirementsPath)
|
||||
if err != nil {
|
||||
as.logger.Warn("failed to read requirements file", "path", requirementsPath, "error", err)
|
||||
} else {
|
||||
as.logger.Info("requirements.txt loaded", "size", len(reqData))
|
||||
}
|
||||
} else {
|
||||
// Fallback: check if requirements.txt exists in the same directory as the algorithm
|
||||
reqPath := filepath.Join(filepath.Dir(algorithmPath), "requirements.txt")
|
||||
if data, err := os.ReadFile(reqPath); err == nil {
|
||||
reqData = data
|
||||
as.logger.Info("found requirements.txt via fallback", "size", len(data))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
as.logger.Info("algorithm loaded", "size", len(algorithmData))
|
||||
as.logger.Info("resource loaded from OCI", "type", resourceType, "size", len(resourceData), "hash_path", hashPath)
|
||||
|
||||
return &DecryptedResource{
|
||||
Data: algorithmData,
|
||||
Data: resourceData,
|
||||
Requirements: reqData,
|
||||
SourceDir: filepath.Dir(algorithmPath),
|
||||
SourceDir: extractDir,
|
||||
}, nil
|
||||
}
|
||||
|
||||
@@ -845,6 +889,8 @@ func (as *agentService) runComputation(state statemachine.State) {
|
||||
as.publishEvent(Starting.String())(state)
|
||||
as.logger.Debug("computation run started")
|
||||
defer func() {
|
||||
as.mu.Lock()
|
||||
defer as.mu.Unlock()
|
||||
if as.runError != nil {
|
||||
as.sm.SendEvent(RunFailed)
|
||||
} else {
|
||||
@@ -852,12 +898,9 @@ func (as *agentService) runComputation(state statemachine.State) {
|
||||
}
|
||||
}()
|
||||
|
||||
if err := os.Mkdir(algorithm.ResultsDir, 0o755); err != nil {
|
||||
as.runError = fmt.Errorf("error creating results directory: %s", err.Error())
|
||||
as.logger.Warn(as.runError.Error())
|
||||
as.publishEvent(Failed.String())(state)
|
||||
return
|
||||
}
|
||||
// Read algo file
|
||||
currentDir, _ := os.Getwd()
|
||||
algoFile := filepath.Join(currentDir, "algo")
|
||||
|
||||
defer func() {
|
||||
if err := os.RemoveAll(algorithm.ResultsDir); err != nil {
|
||||
@@ -866,14 +909,25 @@ func (as *agentService) runComputation(state statemachine.State) {
|
||||
if err := os.RemoveAll(algorithm.DatasetsDir); err != nil {
|
||||
as.logger.Warn(fmt.Sprintf("error removing datasets directory and its contents: %s", err.Error()))
|
||||
}
|
||||
if err := os.Remove(algoFile); err != nil && !os.IsNotExist(err) {
|
||||
as.logger.Warn(fmt.Sprintf("error removing algorithm file: %s", err.Error()))
|
||||
}
|
||||
}()
|
||||
|
||||
// Read algo file
|
||||
currentDir, _ := os.Getwd()
|
||||
algoFile := filepath.Join(currentDir, "algo")
|
||||
if err := os.Mkdir(algorithm.ResultsDir, 0o755); err != nil {
|
||||
as.mu.Lock()
|
||||
as.runError = fmt.Errorf("error creating results directory: %s", err.Error())
|
||||
as.mu.Unlock()
|
||||
as.logger.Warn(as.runError.Error())
|
||||
as.publishEvent(Failed.String())(state)
|
||||
return
|
||||
}
|
||||
|
||||
algoBytes, err := os.ReadFile(algoFile)
|
||||
if err != nil {
|
||||
as.mu.Lock()
|
||||
as.runError = fmt.Errorf("failed to read algo file: %w", err)
|
||||
as.mu.Unlock()
|
||||
as.logger.Warn(as.runError.Error())
|
||||
as.publishEvent(Failed.String())(state)
|
||||
return
|
||||
@@ -891,14 +945,18 @@ func (as *agentService) runComputation(state statemachine.State) {
|
||||
// Datasets implicit on shared FS
|
||||
})
|
||||
if err != nil {
|
||||
as.mu.Lock()
|
||||
as.runError = err
|
||||
as.mu.Unlock()
|
||||
as.logger.Warn(fmt.Sprintf("failed to run computation: %s", err.Error()))
|
||||
as.publishEvent(Failed.String())(state)
|
||||
return
|
||||
}
|
||||
|
||||
if resp.Error != "" {
|
||||
as.mu.Lock()
|
||||
as.runError = errors.New(resp.Error)
|
||||
as.mu.Unlock()
|
||||
as.logger.Warn(fmt.Sprintf("failed to run computation: %s", resp.Error))
|
||||
as.publishEvent(Failed.String())(state)
|
||||
return
|
||||
@@ -906,7 +964,9 @@ func (as *agentService) runComputation(state statemachine.State) {
|
||||
|
||||
results, err := internal.ZipDirectoryToMemory(algorithm.ResultsDir)
|
||||
if err != nil {
|
||||
as.mu.Lock()
|
||||
as.runError = err
|
||||
as.mu.Unlock()
|
||||
as.logger.Warn(fmt.Sprintf("failed to zip results: %s", err.Error()))
|
||||
as.publishEvent(Failed.String())(state)
|
||||
return
|
||||
|
||||
+523
-6
@@ -4,6 +4,8 @@ package agent
|
||||
|
||||
import (
|
||||
"archive/tar"
|
||||
"archive/zip"
|
||||
"bytes"
|
||||
"compress/gzip"
|
||||
"context"
|
||||
"crypto/rand"
|
||||
@@ -46,6 +48,11 @@ func (m *MockOCIClient) PullAndDecrypt(ctx context.Context, source oci.ResourceS
|
||||
return args.Error(0)
|
||||
}
|
||||
|
||||
func (m *MockOCIClient) ToDockerArchive(ctx context.Context, ociDir, destFile string) error {
|
||||
args := m.Called(ctx, ociDir, destFile)
|
||||
return args.Error(0)
|
||||
}
|
||||
|
||||
var (
|
||||
algoPath = "../test/manual/algo/lin_reg.py"
|
||||
reqPath = "../test/manual/algo/requirements.txt"
|
||||
@@ -1097,7 +1104,7 @@ func TestDownloadAlgorithmIfRemote_Success(t *testing.T) {
|
||||
algoContent := []byte("print('hello')")
|
||||
mockOCI.On("PullAndDecrypt", mock.Anything, mock.Anything, mock.Anything).Run(func(args mock.Arguments) {
|
||||
destDir := args.String(2)
|
||||
setupMinimalOCI(t, destDir, "main.py", string(algoContent))
|
||||
setupMinimalOCI(t, destDir, "main.py", algoContent)
|
||||
}).Return(nil)
|
||||
|
||||
svc := newTestAgentService(sm, eventsSvc)
|
||||
@@ -1112,7 +1119,7 @@ func TestDownloadAlgorithmIfRemote_Success(t *testing.T) {
|
||||
AlgoType: "python",
|
||||
Source: &ResourceSource{
|
||||
Type: "oci-image",
|
||||
URL: "docker://test/image",
|
||||
URL: "docker://test/algo-success",
|
||||
},
|
||||
},
|
||||
KBS: KBSConfig{Enabled: true},
|
||||
@@ -1131,7 +1138,57 @@ func TestDownloadAlgorithmIfRemote_Success(t *testing.T) {
|
||||
mockOCI.AssertExpectations(t)
|
||||
}
|
||||
|
||||
func setupMinimalOCI(t *testing.T, ociDir, filename, content string) {
|
||||
func TestDownloadAlgorithmIfRemote_Docker_Success(t *testing.T) {
|
||||
if testing.Short() {
|
||||
t.Skip("skipping in short mode")
|
||||
}
|
||||
|
||||
origDir, _ := os.Getwd()
|
||||
tmpDir := t.TempDir()
|
||||
require.NoError(t, os.Chdir(tmpDir))
|
||||
defer func() { require.NoError(t, os.Chdir(origDir)) }()
|
||||
|
||||
eventsSvc := new(mocks.Service)
|
||||
eventsSvc.On("SendEvent", mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return().Maybe()
|
||||
sm := &smmocks.StateMachine{}
|
||||
sm.On("SendEvent", AlgorithmReceived).Return().Once()
|
||||
|
||||
mockOCI := new(MockOCIClient)
|
||||
mockOCI.On("PullAndDecrypt", mock.Anything, mock.Anything, mock.Anything).Return(nil)
|
||||
|
||||
dummyContent := []byte("dummy docker tar")
|
||||
dummyHash := sha3.Sum256(dummyContent)
|
||||
|
||||
mockOCI.On("ToDockerArchive", mock.Anything, mock.Anything, mock.Anything).Run(func(args mock.Arguments) {
|
||||
destFile := args.String(2)
|
||||
err := os.WriteFile(destFile, dummyContent, 0o644)
|
||||
require.NoError(t, err)
|
||||
}).Return(nil)
|
||||
|
||||
svc := newTestAgentService(sm, eventsSvc)
|
||||
svc.ociClient = mockOCI
|
||||
|
||||
svc.computation = Computation{
|
||||
Algorithm: Algorithm{
|
||||
AlgoType: "docker",
|
||||
Hash: dummyHash,
|
||||
Source: &ResourceSource{
|
||||
Type: "oci-image",
|
||||
URL: "docker://test/algo-docker-success",
|
||||
},
|
||||
},
|
||||
KBS: KBSConfig{Enabled: true},
|
||||
}
|
||||
|
||||
svc.downloadAlgorithmIfRemote(ReceivingAlgorithm)
|
||||
|
||||
assert.Nil(t, svc.runError)
|
||||
assert.True(t, svc.algoReceived)
|
||||
sm.AssertExpectations(t)
|
||||
mockOCI.AssertExpectations(t)
|
||||
}
|
||||
|
||||
func setupMinimalOCI(t *testing.T, ociDir, filename string, content []byte) {
|
||||
t.Helper()
|
||||
blobsDir := filepath.Join(ociDir, "blobs", "sha256")
|
||||
require.NoError(t, os.MkdirAll(blobsDir, 0o755))
|
||||
@@ -1149,7 +1206,8 @@ func setupMinimalOCI(t *testing.T, ociDir, filename, content string) {
|
||||
Size: int64(len(content)),
|
||||
}
|
||||
require.NoError(t, tw.WriteHeader(hdr))
|
||||
_, err = tw.Write([]byte(content))
|
||||
_, err = tw.Write(content)
|
||||
|
||||
require.NoError(t, err)
|
||||
|
||||
require.NoError(t, tw.Close())
|
||||
@@ -1201,7 +1259,7 @@ func TestDownloadDatasetsIfRemote_Success(t *testing.T) {
|
||||
dataContent := []byte("a,b,c\n1,2,3")
|
||||
mockOCI.On("PullAndDecrypt", mock.Anything, mock.Anything, mock.Anything).Run(func(args mock.Arguments) {
|
||||
destDir := args.String(2)
|
||||
setupMinimalOCI(t, destDir, "data.csv", string(dataContent))
|
||||
setupMinimalOCI(t, destDir, "data.csv", dataContent)
|
||||
}).Return(nil)
|
||||
|
||||
svc := newTestAgentService(sm, eventsSvc)
|
||||
@@ -1217,7 +1275,7 @@ func TestDownloadDatasetsIfRemote_Success(t *testing.T) {
|
||||
Hash: dataHash,
|
||||
Source: &ResourceSource{
|
||||
Type: "oci-image",
|
||||
URL: "docker://test/image",
|
||||
URL: "docker://test/data-success",
|
||||
},
|
||||
},
|
||||
},
|
||||
@@ -1234,3 +1292,462 @@ func TestDownloadDatasetsIfRemote_Success(t *testing.T) {
|
||||
sm.AssertExpectations(t)
|
||||
mockOCI.AssertExpectations(t)
|
||||
}
|
||||
|
||||
func TestDownloadDatasetsIfRemote_Decompress(t *testing.T) {
|
||||
if testing.Short() {
|
||||
t.Skip("skipping in short mode")
|
||||
}
|
||||
|
||||
origDir, _ := os.Getwd()
|
||||
tmpDir := t.TempDir()
|
||||
require.NoError(t, os.Chdir(tmpDir))
|
||||
defer func() { require.NoError(t, os.Chdir(origDir)) }()
|
||||
|
||||
eventsSvc := new(mocks.Service)
|
||||
eventsSvc.On("SendEvent", mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return().Maybe()
|
||||
sm := &smmocks.StateMachine{}
|
||||
sm.On("SendEvent", DataReceived).Return().Maybe()
|
||||
sm.On("SendEvent", RunFailed).Return().Maybe()
|
||||
|
||||
mockOCI := new(MockOCIClient)
|
||||
|
||||
// Create a zip file in memory
|
||||
var buf bytes.Buffer
|
||||
zw := zip.NewWriter(&buf)
|
||||
f, err := zw.Create("test.txt")
|
||||
require.NoError(t, err)
|
||||
_, err = f.Write([]byte("hello zip"))
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, zw.Close())
|
||||
zipData := buf.Bytes()
|
||||
|
||||
mockOCI.On("PullAndDecrypt", mock.Anything, mock.Anything, mock.Anything).Run(func(args mock.Arguments) {
|
||||
destDir := args.String(2)
|
||||
setupMinimalOCI(t, destDir, "data.zip", zipData)
|
||||
}).Return(nil)
|
||||
|
||||
svc := newTestAgentService(sm, eventsSvc)
|
||||
svc.ociClient = mockOCI
|
||||
|
||||
dataHash := sha3.Sum256(zipData)
|
||||
|
||||
svc.computation = Computation{
|
||||
Datasets: []Dataset{
|
||||
{
|
||||
Filename: "data.zip",
|
||||
Hash: dataHash,
|
||||
Decompress: true,
|
||||
Source: &ResourceSource{
|
||||
Type: "oci-image",
|
||||
URL: "docker://test/data-decompress",
|
||||
},
|
||||
},
|
||||
},
|
||||
KBS: KBSConfig{Enabled: true},
|
||||
}
|
||||
|
||||
err = os.MkdirAll(algorithm.DatasetsDir, 0o755)
|
||||
require.NoError(t, err)
|
||||
|
||||
svc.downloadDatasetsIfRemote(ReceivingData)
|
||||
|
||||
assert.Nil(t, svc.runError)
|
||||
assert.Len(t, svc.computation.Datasets, 0)
|
||||
// Check if file was decompressed
|
||||
decompressedFile := filepath.Join(algorithm.DatasetsDir, "test.txt")
|
||||
_, err = os.Stat(decompressedFile)
|
||||
assert.NoError(t, err)
|
||||
|
||||
sm.AssertExpectations(t)
|
||||
mockOCI.AssertExpectations(t)
|
||||
}
|
||||
|
||||
func TestDownloadAlgorithmIfRemote_ErrorPathsInternal(t *testing.T) {
|
||||
origDir, _ := os.Getwd()
|
||||
tmpDir := t.TempDir()
|
||||
require.NoError(t, os.Chdir(tmpDir))
|
||||
defer func() { require.NoError(t, os.Chdir(origDir)) }()
|
||||
|
||||
eventsSvc := new(mocks.Service)
|
||||
eventsSvc.EXPECT().SendEvent(mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return().Maybe()
|
||||
|
||||
t.Run("hash mismatch", func(t *testing.T) {
|
||||
sm := &smmocks.StateMachine{}
|
||||
sm.On("SendEvent", RunFailed).Return().Once()
|
||||
|
||||
mockOCI := new(MockOCIClient)
|
||||
mockOCI.On("PullAndDecrypt", mock.Anything, mock.Anything, mock.Anything).Run(func(args mock.Arguments) {
|
||||
destDir := args.String(2)
|
||||
setupMinimalOCI(t, destDir, "main.py", []byte("wrong content"))
|
||||
}).Return(nil)
|
||||
|
||||
svc := newTestAgentService(sm, eventsSvc)
|
||||
svc.ociClient = mockOCI
|
||||
|
||||
svc.computation = Computation{
|
||||
Algorithm: Algorithm{
|
||||
Hash: sha3.Sum256([]byte("expected content")),
|
||||
AlgoType: "python",
|
||||
Source: &ResourceSource{
|
||||
Type: "oci-image",
|
||||
URL: "docker://test/algo-hash-mismatch",
|
||||
},
|
||||
},
|
||||
KBS: KBSConfig{Enabled: true},
|
||||
}
|
||||
|
||||
svc.downloadAlgorithmIfRemote(ReceivingAlgorithm)
|
||||
assert.Error(t, svc.runError)
|
||||
assert.Contains(t, svc.runError.Error(), "algorithm hash mismatch")
|
||||
sm.AssertExpectations(t)
|
||||
})
|
||||
|
||||
t.Run("create algo file failure", func(t *testing.T) {
|
||||
sm := &smmocks.StateMachine{}
|
||||
sm.On("SendEvent", RunFailed).Return().Once()
|
||||
|
||||
// Create a directory named "algo" to make file creation fail
|
||||
require.NoError(t, os.Mkdir("algo", 0o755))
|
||||
defer os.RemoveAll("algo")
|
||||
|
||||
mockOCI := new(MockOCIClient)
|
||||
algoContent := "print(1)"
|
||||
mockOCI.On("PullAndDecrypt", mock.Anything, mock.Anything, mock.Anything).Run(func(args mock.Arguments) {
|
||||
destDir := args.String(2)
|
||||
setupMinimalOCI(t, destDir, "main.py", []byte(algoContent))
|
||||
}).Return(nil)
|
||||
|
||||
svc := newTestAgentService(sm, eventsSvc)
|
||||
svc.ociClient = mockOCI
|
||||
|
||||
svc.computation = Computation{
|
||||
Algorithm: Algorithm{
|
||||
Hash: sha3.Sum256([]byte(algoContent)),
|
||||
AlgoType: "python",
|
||||
Source: &ResourceSource{
|
||||
Type: "oci-image",
|
||||
URL: "docker://test/algo-create-fail",
|
||||
},
|
||||
},
|
||||
KBS: KBSConfig{Enabled: true},
|
||||
}
|
||||
|
||||
svc.downloadAlgorithmIfRemote(ReceivingAlgorithm)
|
||||
assert.Error(t, svc.runError)
|
||||
assert.Contains(t, svc.runError.Error(), "error creating algorithm file")
|
||||
sm.AssertExpectations(t)
|
||||
})
|
||||
t.Run("extraction failure", func(t *testing.T) {
|
||||
sm := &smmocks.StateMachine{}
|
||||
sm.On("SendEvent", RunFailed).Return().Once()
|
||||
|
||||
mockOCI := new(MockOCIClient)
|
||||
mockOCI.On("PullAndDecrypt", mock.Anything, mock.Anything, mock.Anything).Run(func(args mock.Arguments) {
|
||||
destDir := args.String(2)
|
||||
// Setup OCI with NO main.py or any algorithm file
|
||||
require.NoError(t, os.MkdirAll(filepath.Join(destDir, "blobs"), 0o755))
|
||||
// Create a legit-looking but empty index.json
|
||||
require.NoError(t, os.WriteFile(filepath.Join(destDir, "index.json"), []byte(`{"schemaVersion":2,"manifests":[]}`), 0o644))
|
||||
}).Return(nil)
|
||||
|
||||
svc := newTestAgentService(sm, eventsSvc)
|
||||
svc.ociClient = mockOCI
|
||||
|
||||
svc.computation = Computation{
|
||||
Algorithm: Algorithm{
|
||||
AlgoType: "python",
|
||||
Source: &ResourceSource{
|
||||
Type: "oci-image",
|
||||
URL: "docker://test/image",
|
||||
},
|
||||
},
|
||||
KBS: KBSConfig{Enabled: true},
|
||||
}
|
||||
|
||||
svc.downloadAlgorithmIfRemote(ReceivingAlgorithm)
|
||||
assert.Error(t, svc.runError)
|
||||
assert.Contains(t, svc.runError.Error(), "no manifests found")
|
||||
sm.AssertExpectations(t)
|
||||
})
|
||||
}
|
||||
|
||||
func TestDownloadDatasetsIfRemote_ErrorPathsInternal(t *testing.T) {
|
||||
origDir, _ := os.Getwd()
|
||||
tmpDir := t.TempDir()
|
||||
require.NoError(t, os.Chdir(tmpDir))
|
||||
defer func() { require.NoError(t, os.Chdir(origDir)) }()
|
||||
|
||||
// Use a fresh mock in each subtest to avoid state pollution
|
||||
|
||||
t.Run("dataset create file failure", func(t *testing.T) {
|
||||
eventsSvc := mocks.NewService(t)
|
||||
eventsSvc.On("SendEvent", mock.Anything, mock.Anything, mock.Anything, mock.MatchedBy(func(json.RawMessage) bool { return true })).Return().Maybe()
|
||||
sm := &smmocks.StateMachine{}
|
||||
sm.On("SendEvent", RunFailed).Return().Once()
|
||||
|
||||
// Create a directory named "data.csv" in datasets dir to make file creation fail
|
||||
require.NoError(t, os.MkdirAll(filepath.Join(algorithm.DatasetsDir, "data.csv"), 0o755))
|
||||
defer os.RemoveAll(algorithm.DatasetsDir)
|
||||
|
||||
mockOCI := new(MockOCIClient)
|
||||
dataContent := "a,b,c"
|
||||
mockOCI.On("PullAndDecrypt", mock.Anything, mock.Anything, mock.Anything).Run(func(args mock.Arguments) {
|
||||
destDir := args.String(2)
|
||||
setupMinimalOCI(t, destDir, "data.csv", []byte(dataContent))
|
||||
}).Return(nil)
|
||||
|
||||
svc := newTestAgentService(sm, eventsSvc)
|
||||
svc.ociClient = mockOCI
|
||||
|
||||
svc.computation = Computation{
|
||||
Datasets: []Dataset{
|
||||
{
|
||||
Filename: "data.csv",
|
||||
Hash: sha3.Sum256([]byte(dataContent)),
|
||||
Source: &ResourceSource{
|
||||
Type: "oci-image",
|
||||
URL: "docker://test/data-create-fail",
|
||||
},
|
||||
},
|
||||
},
|
||||
KBS: KBSConfig{Enabled: true},
|
||||
}
|
||||
|
||||
svc.downloadDatasetsIfRemote(ReceivingData)
|
||||
sm.AssertExpectations(t)
|
||||
})
|
||||
|
||||
t.Run("dataset hash mismatch", func(t *testing.T) {
|
||||
eventsSvc := mocks.NewService(t)
|
||||
eventsSvc.On("SendEvent", mock.Anything, mock.Anything, mock.Anything, mock.MatchedBy(func(json.RawMessage) bool { return true })).Return().Maybe()
|
||||
origDir, _ := os.Getwd()
|
||||
tmpDir := t.TempDir()
|
||||
require.NoError(t, os.Chdir(tmpDir))
|
||||
defer func() { _ = os.Chdir(origDir) }()
|
||||
|
||||
sm := &smmocks.StateMachine{}
|
||||
sm.On("SendEvent", RunFailed).Return().Once()
|
||||
|
||||
mockOCI := new(MockOCIClient)
|
||||
dataContent := "wrong content"
|
||||
mockOCI.On("PullAndDecrypt", mock.Anything, mock.Anything, mock.Anything).Run(func(args mock.Arguments) {
|
||||
destDir := args.String(2)
|
||||
setupMinimalOCI(t, destDir, "data.csv", []byte(dataContent))
|
||||
}).Return(nil)
|
||||
|
||||
svc := newTestAgentService(sm, eventsSvc)
|
||||
svc.ociClient = mockOCI
|
||||
|
||||
svc.computation = Computation{
|
||||
Datasets: []Dataset{
|
||||
{
|
||||
Filename: "data.csv",
|
||||
Hash: sha3.Sum256([]byte("expected content")),
|
||||
Source: &ResourceSource{
|
||||
Type: "oci-image",
|
||||
URL: "docker://test/data-mismatch",
|
||||
},
|
||||
},
|
||||
},
|
||||
KBS: KBSConfig{Enabled: true},
|
||||
}
|
||||
|
||||
err := os.MkdirAll(algorithm.DatasetsDir, 0o755)
|
||||
require.NoError(t, err)
|
||||
|
||||
svc.downloadDatasetsIfRemote(ReceivingData)
|
||||
if svc.runError == nil {
|
||||
t.Fatalf("runError should not be nil in hash mismatch test")
|
||||
}
|
||||
assert.Contains(t, svc.runError.Error(), "dataset data.csv hash mismatch")
|
||||
sm.AssertExpectations(t)
|
||||
})
|
||||
|
||||
t.Run("dataset unzip failure", func(t *testing.T) {
|
||||
eventsSvc := mocks.NewService(t)
|
||||
eventsSvc.On("SendEvent", mock.Anything, mock.Anything, mock.Anything, mock.MatchedBy(func(json.RawMessage) bool { return true })).Return().Maybe()
|
||||
origDir, _ := os.Getwd()
|
||||
tmpDir := t.TempDir()
|
||||
require.NoError(t, os.Chdir(tmpDir))
|
||||
defer func() { _ = os.Chdir(origDir) }()
|
||||
|
||||
sm := &smmocks.StateMachine{}
|
||||
sm.On("SendEvent", RunFailed).Return().Once()
|
||||
|
||||
mockOCI := new(MockOCIClient)
|
||||
// Provide invalid zip content
|
||||
dataContent := "not a zip file"
|
||||
mockOCI.On("PullAndDecrypt", mock.Anything, mock.Anything, mock.Anything).Run(func(args mock.Arguments) {
|
||||
destDir := args.String(2)
|
||||
setupMinimalOCI(t, destDir, "data.zip", []byte(dataContent))
|
||||
}).Return(nil)
|
||||
|
||||
svc := newTestAgentService(sm, eventsSvc)
|
||||
svc.ociClient = mockOCI
|
||||
|
||||
svc.computation = Computation{
|
||||
Datasets: []Dataset{
|
||||
{
|
||||
Filename: "data.zip",
|
||||
Hash: sha3.Sum256([]byte(dataContent)),
|
||||
Decompress: true,
|
||||
Source: &ResourceSource{
|
||||
Type: "oci-image",
|
||||
URL: "docker://test/data-unzip-fail",
|
||||
},
|
||||
},
|
||||
},
|
||||
KBS: KBSConfig{Enabled: true},
|
||||
}
|
||||
|
||||
err := os.MkdirAll(algorithm.DatasetsDir, 0o755)
|
||||
require.NoError(t, err)
|
||||
|
||||
svc.downloadDatasetsIfRemote(ReceivingData)
|
||||
if svc.runError == nil {
|
||||
t.Fatalf("runError should not be nil in unzip failure test")
|
||||
}
|
||||
assert.Contains(t, svc.runError.Error(), "failed to unzip dataset")
|
||||
sm.AssertExpectations(t)
|
||||
})
|
||||
}
|
||||
|
||||
func TestAlgo_RemoteSource(t *testing.T) {
|
||||
if testing.Short() {
|
||||
t.Skip("skipping in short mode")
|
||||
}
|
||||
|
||||
origDir, _ := os.Getwd()
|
||||
tmpDir := t.TempDir()
|
||||
require.NoError(t, os.Chdir(tmpDir))
|
||||
defer func() { require.NoError(t, os.Chdir(origDir)) }()
|
||||
|
||||
eventsSvc := new(mocks.Service)
|
||||
eventsSvc.EXPECT().SendEvent(mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return().Maybe()
|
||||
sm := &smmocks.StateMachine{}
|
||||
sm.On("GetState").Return(ReceivingAlgorithm)
|
||||
sm.On("SendEvent", AlgorithmReceived).Return().Once()
|
||||
|
||||
mockOCI := new(MockOCIClient)
|
||||
algoContent := []byte("print('remote algo')")
|
||||
algoHash := sha3.Sum256(algoContent)
|
||||
|
||||
mockOCI.On("PullAndDecrypt", mock.Anything, mock.Anything, mock.Anything).Run(func(args mock.Arguments) {
|
||||
destDir := args.String(2)
|
||||
setupMinimalOCI(t, destDir, "main.py", algoContent)
|
||||
}).Return(nil)
|
||||
|
||||
svc := &agentService{
|
||||
logger: slog.Default(),
|
||||
eventSvc: eventsSvc,
|
||||
sm: sm,
|
||||
ociClient: mockOCI,
|
||||
computation: Computation{
|
||||
Algorithm: Algorithm{
|
||||
Hash: algoHash,
|
||||
AlgoType: "python",
|
||||
Source: &ResourceSource{
|
||||
Type: "oci-image",
|
||||
URL: "docker://test/algo-remote",
|
||||
},
|
||||
},
|
||||
KBS: KBSConfig{Enabled: true},
|
||||
},
|
||||
}
|
||||
|
||||
ctx := metadata.NewIncomingContext(context.Background(), metadata.Pairs(algorithm.AlgoTypeKey, "python"))
|
||||
err := svc.Algo(ctx, Algorithm{})
|
||||
assert.NoError(t, err)
|
||||
assert.True(t, svc.algoReceived)
|
||||
sm.AssertExpectations(t)
|
||||
mockOCI.AssertExpectations(t)
|
||||
}
|
||||
|
||||
func TestData_RemoteSource(t *testing.T) {
|
||||
if testing.Short() {
|
||||
t.Skip("skipping in short mode")
|
||||
}
|
||||
|
||||
origDir, _ := os.Getwd()
|
||||
tmpDir := t.TempDir()
|
||||
require.NoError(t, os.Chdir(tmpDir))
|
||||
defer func() { require.NoError(t, os.Chdir(origDir)) }()
|
||||
|
||||
eventsSvc := new(mocks.Service)
|
||||
eventsSvc.EXPECT().SendEvent(mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return().Maybe()
|
||||
sm := &smmocks.StateMachine{}
|
||||
sm.On("GetState").Return(ReceivingData)
|
||||
sm.On("SendEvent", DataReceived).Return().Once()
|
||||
|
||||
mockOCI := new(MockOCIClient)
|
||||
dataContent := []byte("remote data")
|
||||
dataHash := sha3.Sum256(dataContent)
|
||||
|
||||
mockOCI.On("PullAndDecrypt", mock.Anything, mock.Anything, mock.Anything).Run(func(args mock.Arguments) {
|
||||
destDir := args.String(2)
|
||||
setupMinimalOCI(t, destDir, "data.csv", dataContent)
|
||||
}).Return(nil)
|
||||
|
||||
svc := &agentService{
|
||||
logger: slog.Default(),
|
||||
eventSvc: eventsSvc,
|
||||
sm: sm,
|
||||
ociClient: mockOCI,
|
||||
computation: Computation{
|
||||
Datasets: []Dataset{
|
||||
{
|
||||
Filename: "data.csv",
|
||||
Hash: dataHash,
|
||||
Source: &ResourceSource{
|
||||
Type: "oci-image",
|
||||
URL: "docker://test/data-remote",
|
||||
},
|
||||
},
|
||||
},
|
||||
KBS: KBSConfig{Enabled: true},
|
||||
},
|
||||
}
|
||||
|
||||
err := os.MkdirAll(algorithm.DatasetsDir, 0o755)
|
||||
require.NoError(t, err)
|
||||
|
||||
ctx := context.Background()
|
||||
err = svc.Data(ctx, Dataset{})
|
||||
assert.NoError(t, err)
|
||||
assert.Len(t, svc.computation.Datasets, 0)
|
||||
sm.AssertExpectations(t)
|
||||
mockOCI.AssertExpectations(t)
|
||||
}
|
||||
|
||||
func TestRunComputation_Success(t *testing.T) {
|
||||
origDir, _ := os.Getwd()
|
||||
tmpDir := t.TempDir()
|
||||
require.NoError(t, os.Chdir(tmpDir))
|
||||
defer func() { require.NoError(t, os.Chdir(origDir)) }()
|
||||
|
||||
// Write a dummy algo file
|
||||
require.NoError(t, os.WriteFile("algo", []byte("#!/bin/sh\necho ok\n"), 0o755))
|
||||
|
||||
runnerCli := new(runnermocks.Client)
|
||||
runnerCli.On("Run", mock.Anything, mock.Anything).Return(&runnerpb.RunResponse{}, nil)
|
||||
|
||||
eventsSvc := new(mocks.Service)
|
||||
eventsSvc.EXPECT().SendEvent(mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return().Maybe()
|
||||
|
||||
sm := &smmocks.StateMachine{}
|
||||
sm.On("SendEvent", RunComplete).Return().Once()
|
||||
|
||||
svc := &agentService{
|
||||
logger: slog.Default(),
|
||||
eventSvc: eventsSvc,
|
||||
sm: sm,
|
||||
runnerClient: runnerCli,
|
||||
computation: Computation{ID: "test-run"},
|
||||
}
|
||||
|
||||
svc.runComputation(Running)
|
||||
|
||||
assert.Nil(t, svc.runError)
|
||||
sm.AssertExpectations(t)
|
||||
runnerCli.AssertExpectations(t)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user