diff --git a/agent/auth/auth.go b/agent/auth/auth.go index f94264d5..0a1f5985 100644 --- a/agent/auth/auth.go +++ b/agent/auth/auth.go @@ -146,9 +146,9 @@ func (s *service) AuthenticateUser(ctx context.Context, role UserRole) (context. } } case DataProviderRole: - for i, dp := range s.datasetProviders { + for _, dp := range s.datasetProviders { if err := verifySignature(role, signature, dp); err == nil { - return agent.IndexToContext(ctx, i), nil + return ctx, nil } } case AlgorithmProviderRole: diff --git a/agent/auth/auth_test.go b/agent/auth/auth_test.go index e04bd18b..dae13c37 100644 --- a/agent/auth/auth_test.go +++ b/agent/auth/auth_test.go @@ -124,7 +124,7 @@ func TestAuthenticateUser(t *testing.T) { if err == nil { switch id, ok := agent.IndexFromContext(ctx); { - case tc.role == ConsumerRole, tc.role == DataProviderRole: + case tc.role == ConsumerRole: assert.True(t, ok, "expected index in context") assert.Equal(t, 0, id, "expected index 0 in context") default: diff --git a/agent/service.go b/agent/service.go index de3d8d0d..000128ff 100644 --- a/agent/service.go +++ b/agent/service.go @@ -179,33 +179,37 @@ func (as *agentService) Data(ctx context.Context, dataset Dataset) error { hash := sha3.Sum256(dataset.Dataset) - index, ok := IndexFromContext(ctx) - if !ok { + matched := false + for i, d := range as.computation.Datasets { + if hash == d.Hash { + if d.Filename != "" && d.Filename != dataset.Filename { + return ErrFileNameMismatch + } + + as.computation.Datasets = slices.Delete(as.computation.Datasets, i, i+1) + + f, err := os.Create(fmt.Sprintf("%s/%s", algorithm.DatasetsDir, dataset.Filename)) + 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) + } + + matched = true + break + } + } + + if !matched { return ErrUndeclaredDataset } - if hash != as.computation.Datasets[index].Hash { - return ErrHashMismatch - } - - if as.computation.Datasets[index].Filename != "" && as.computation.Datasets[index].Filename != dataset.Filename { - return ErrFileNameMismatch - } - - as.computation.Datasets = slices.Delete(as.computation.Datasets, index, index+1) - - f, err := os.Create(fmt.Sprintf("%s/%s", algorithm.DatasetsDir, dataset.Filename)) - 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) - } - + // Check if all datasets have been received if len(as.computation.Datasets) == 0 { as.sm.SendEvent(dataReceived) }