package storage import ( "bytes" "context" "errors" "io" "strings" "sync" "testing" "github.com/aws/aws-sdk-go-v2/service/s3" "github.com/aws/aws-sdk-go-v2/service/s3/types" ) type mockS3Client struct { mu sync.Mutex objects map[string][]byte } func (m *mockS3Client) PutObject(_ context.Context, params *s3.PutObjectInput, _ ...func(*s3.Options)) (*s3.PutObjectOutput, error) { m.mu.Lock() defer m.mu.Unlock() data, err := io.ReadAll(params.Body) if err != nil { return nil, err } m.objects[*params.Bucket+"/"+*params.Key] = data return &s3.PutObjectOutput{}, nil } func (m *mockS3Client) GetObject(_ context.Context, params *s3.GetObjectInput, _ ...func(*s3.Options)) (*s3.GetObjectOutput, error) { m.mu.Lock() defer m.mu.Unlock() data, ok := m.objects[*params.Bucket+"/"+*params.Key] if !ok { return nil, &types.NoSuchKey{} } return &s3.GetObjectOutput{Body: io.NopCloser(bytes.NewReader(data))}, nil } func (m *mockS3Client) DeleteObject(_ context.Context, params *s3.DeleteObjectInput, _ ...func(*s3.Options)) (*s3.DeleteObjectOutput, error) { m.mu.Lock() defer m.mu.Unlock() delete(m.objects, *params.Bucket+"/"+*params.Key) return &s3.DeleteObjectOutput{}, nil } func TestValidateKey(t *testing.T) { valid := []string{ "did:plc:xyz123/go-mod-v1", "did:web:spindle.example.com/cache.tar", "abc", "a/b/c/d", } for _, k := range valid { if err := ValidateKey(k); err != nil { t.Errorf("ValidateKey(%q) = %v, want nil", k, err) } } invalid := []string{ "", "../escape", "a/../../b", "/leading", "trailing/", "double//slash", "with space", "with\\backslash", "-leading-dash", } for _, k := range invalid { if err := ValidateKey(k); err == nil { t.Errorf("ValidateKey(%q) = nil, want error", k) } } } func TestDiskRoundTrip(t *testing.T) { ctx := context.Background() d, err := NewDisk(t.TempDir()) if err != nil { t.Fatal(err) } key := "did:plc:xyz/go-mod-v1" if err := d.Put(ctx, key, strings.NewReader("archive-bytes")); err != nil { t.Fatal(err) } rc, err := d.Get(ctx, key) if err != nil { t.Fatal(err) } got, err := io.ReadAll(rc) rc.Close() if err != nil { t.Fatal(err) } if string(got) != "archive-bytes" { t.Fatalf("got %q, want %q", got, "archive-bytes") } // overwrite if err := d.Put(ctx, key, strings.NewReader("new-bytes")); err != nil { t.Fatal(err) } rc, _ = d.Get(ctx, key) got, _ = io.ReadAll(rc) rc.Close() if string(got) != "new-bytes" { t.Fatalf("overwrite: got %q, want %q", got, "new-bytes") } if err := d.Delete(ctx, key); err != nil { t.Fatalf("delete: %v", err) } if _, err := d.Get(ctx, key); !errors.Is(err, ErrNotExist) { t.Fatalf("get deleted: got %v, want ErrNotExist", err) } if err := d.Delete(ctx, key); err != nil { t.Fatalf("delete missing: %v", err) } } func TestDiskRejectsTraversal(t *testing.T) { ctx := context.Background() d, err := NewDisk(t.TempDir()) if err != nil { t.Fatal(err) } if err := d.Put(ctx, "../evil", strings.NewReader("x")); err == nil { t.Fatal("put with traversal key succeeded") } if _, err := d.Get(ctx, "../evil"); err == nil { t.Fatal("get with traversal key succeeded") } } func TestS3RoundTrip(t *testing.T) { ctx := context.Background() client := &mockS3Client{objects: make(map[string][]byte)} store, err := NewS3(client, "bucket", "prefix") if err != nil { t.Fatal(err) } if err := store.Put(ctx, "logs/test.log", strings.NewReader("artifact")); err != nil { t.Fatal(err) } rc, err := store.Get(ctx, "logs/test.log") if err != nil { t.Fatal(err) } got, err := io.ReadAll(rc) rc.Close() if err != nil { t.Fatal(err) } if string(got) != "artifact" { t.Fatalf("content = %q", got) } if err := store.Delete(ctx, "logs/test.log"); err != nil { t.Fatal(err) } if _, err := store.Get(ctx, "logs/test.log"); !errors.Is(err, ErrNotExist) { t.Fatalf("get deleted = %v, want ErrNotExist", err) } }