Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions go.mod
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
35 changes: 27 additions & 8 deletions internal/agent/files.go
Original file line number Diff line number Diff line change
Expand Up @@ -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) {
Expand Down
128 changes: 128 additions & 0 deletions internal/agent/files_test.go
Original file line number Diff line number Diff line change
@@ -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)
}
}
}
Loading