diff --git a/go.mod b/go.mod index a28ae8f..7b75538 100644 --- a/go.mod +++ b/go.mod @@ -60,6 +60,7 @@ require ( golang.org/x/crypto v0.54.0 // indirect golang.org/x/exp v0.0.0-20250305212735-054e65f0b394 // indirect golang.org/x/net v0.57.0 // indirect + golang.org/x/sync v0.22.0 // indirect golang.org/x/sys v0.47.0 // indirect golang.org/x/term v0.45.0 // indirect golang.org/x/text v0.40.0 // indirect diff --git a/internal/agent/files.go b/internal/agent/files.go index 03f8a70..ea6de8f 100644 --- a/internal/agent/files.go +++ b/internal/agent/files.go @@ -8,21 +8,40 @@ import ( "time" "github.com/google/research-cli/internal/utils" + "golang.org/x/sync/errgroup" "google.golang.org/genai" ) func (a *ResearchAgent) UploadFiles(ctx context.Context, filePaths []string) ([]string, error) { - var uris []string - for _, path := range filePaths { - uri, err := a.uploadFile(ctx, path) - if err != nil { - return nil, fmt.Errorf("failed to upload %s: %w", path, err) - } + if len(filePaths) == 0 { + return nil, nil + } + + g, ctx := errgroup.WithContext(ctx) + uris := make([]string, len(filePaths)) + + for i, path := range filePaths { + g.Go(func() error { + uri, err := a.uploadFile(ctx, path) + if err != nil { + return fmt.Errorf("failed to upload %s: %w", path, err) + } + uris[i] = uri + return nil + }) + } + + if err := g.Wait(); err != nil { + return nil, err + } + + result := make([]string, 0, len(uris)) + for _, uri := range uris { if uri != "" { - uris = append(uris, uri) + result = append(result, uri) } } - return uris, nil + return result, nil } func (a *ResearchAgent) uploadFile(ctx context.Context, path string) (string, error) { diff --git a/internal/agent/files_test.go b/internal/agent/files_test.go new file mode 100644 index 0000000..5d4e112 --- /dev/null +++ b/internal/agent/files_test.go @@ -0,0 +1,128 @@ +package agent + +import ( + "context" + "fmt" + "reflect" + "testing" + "time" + + "golang.org/x/sync/errgroup" +) + +type uploadFileFunc func(ctx context.Context, path string) (string, error) + +func uploadFilesSequential(ctx context.Context, filePaths []string, uploadFn uploadFileFunc) ([]string, error) { + var uris []string + for _, path := range filePaths { + uri, err := uploadFn(ctx, path) + if err != nil { + return nil, fmt.Errorf("failed to upload %s: %w", path, err) + } + if uri != "" { + uris = append(uris, uri) + } + } + return uris, nil +} + +func uploadFilesConcurrent(ctx context.Context, filePaths []string, uploadFn uploadFileFunc) ([]string, error) { + if len(filePaths) == 0 { + return nil, nil + } + + g, ctx := errgroup.WithContext(ctx) + uris := make([]string, len(filePaths)) + + for i, path := range filePaths { + g.Go(func() error { + uri, err := uploadFn(ctx, path) + if err != nil { + return fmt.Errorf("failed to upload %s: %w", path, err) + } + uris[i] = uri + return nil + }) + } + + if err := g.Wait(); err != nil { + return nil, err + } + + result := make([]string, 0, len(uris)) + for _, uri := range uris { + if uri != "" { + result = append(result, uri) + } + } + return result, nil +} + +func TestUploadFilesConcurrentPreservesOrder(t *testing.T) { + paths := []string{"a.txt", "b.txt", "c.txt", "d.txt"} + mockUpload := func(ctx context.Context, path string) (string, error) { + if path == "a.txt" { + time.Sleep(30 * time.Millisecond) + } else { + time.Sleep(5 * time.Millisecond) + } + return "uri://" + path, nil + } + + uris, err := uploadFilesConcurrent(t.Context(), paths, mockUpload) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + want := []string{"uri://a.txt", "uri://b.txt", "uri://c.txt", "uri://d.txt"} + if !reflect.DeepEqual(uris, want) { + t.Fatalf("got uris %v, want %v", uris, want) + } +} + +func TestUploadFilesConcurrentErrorHandling(t *testing.T) { + paths := []string{"a.txt", "fail.txt", "c.txt"} + mockUpload := func(ctx context.Context, path string) (string, error) { + if path == "fail.txt" { + return "", fmt.Errorf("upload error") + } + return "uri://" + path, nil + } + + _, err := uploadFilesConcurrent(t.Context(), paths, mockUpload) + if err == nil { + t.Fatal("expected error, got nil") + } +} + +func BenchmarkUploadFiles_Sequential(b *testing.B) { + paths := []string{"file1.txt", "file2.txt", "file3.txt", "file4.txt", "file5.txt"} + mockUpload := func(ctx context.Context, path string) (string, error) { + time.Sleep(10 * time.Millisecond) + return "uri://" + path, nil + } + + b.ResetTimer() + for i := 0; i < b.N; i++ { + _, err := uploadFilesSequential(context.Background(), paths, mockUpload) + if err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkUploadFiles_Concurrent(b *testing.B) { + paths := []string{"file1.txt", "file2.txt", "file3.txt", "file4.txt", "file5.txt"} + mockUpload := func(ctx context.Context, path string) (string, error) { + time.Sleep(10 * time.Millisecond) + return "uri://" + path, nil + } + + b.ResetTimer() + for i := 0; i < b.N; i++ { + _, err := uploadFilesConcurrent(context.Background(), paths, mockUpload) + if err != nil { + b.Fatal(err) + } + } +}