Files
zw-go/client_test.go

1154 lines
33 KiB
Go

package zw
import (
"bytes"
"context"
"errors"
"fmt"
"io"
"io/fs"
"net"
"path/filepath"
"sync/atomic"
"testing"
"time"
zwdaemonv1 "git.pablu.de/pablu/zw-go/gen/zwdaemon/v1"
"google.golang.org/grpc"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/status"
"google.golang.org/protobuf/proto"
"google.golang.org/protobuf/types/known/anypb"
)
const testZwErrorDetailTypeURL = "type.googleapis.com/zw.daemon.v1.ZwErrorDetail"
type testService struct {
zwdaemonv1.UnimplementedZwDaemonServiceServer
response *zwdaemonv1.GetServerInfoResponse
err error
createRequests chan *zwdaemonv1.CreateRequest
createErr error
unsealRequests chan *zwdaemonv1.UnsealRequest
unsealErr error
sealCalls chan struct{}
sealErr error
listFilesRequests chan *zwdaemonv1.ListFilesRequest
listFilesResponse *zwdaemonv1.ListFilesResponse
listFilesErr error
deleteFileRequests chan *zwdaemonv1.DeleteFileRequest
deleteFileErr error
moveFileRequests chan *zwdaemonv1.MoveFileRequest
moveFileErr error
copyFileRequests chan *zwdaemonv1.CopyFileRequest
copyFileErr error
compactCalls chan struct{}
compactErr error
getFile func(*zwdaemonv1.GetFileRequest, grpc.ServerStreamingServer[zwdaemonv1.GetFileResponse]) error
getFileRange func(*zwdaemonv1.GetFileRangeRequest, grpc.ServerStreamingServer[zwdaemonv1.GetFileRangeResponse]) error
addFile func(grpc.ClientStreamingServer[zwdaemonv1.AddFileRequest, zwdaemonv1.AddFileResponse]) error
replaceFile func(grpc.ClientStreamingServer[zwdaemonv1.ReplaceFileRequest, zwdaemonv1.ReplaceFileResponse]) error
}
func (s *testService) GetServerInfo(
context.Context,
*zwdaemonv1.GetServerInfoRequest,
) (*zwdaemonv1.GetServerInfoResponse, error) {
return s.response, s.err
}
func (s *testService) Create(
_ context.Context,
request *zwdaemonv1.CreateRequest,
) (*zwdaemonv1.CreateResponse, error) {
if s.createRequests != nil {
s.createRequests <- request
}
return &zwdaemonv1.CreateResponse{}, s.createErr
}
func (s *testService) Unseal(
_ context.Context,
request *zwdaemonv1.UnsealRequest,
) (*zwdaemonv1.UnsealResponse, error) {
if s.unsealRequests != nil {
s.unsealRequests <- request
}
return &zwdaemonv1.UnsealResponse{}, s.unsealErr
}
func (s *testService) Seal(
context.Context,
*zwdaemonv1.SealRequest,
) (*zwdaemonv1.SealResponse, error) {
if s.sealCalls != nil {
s.sealCalls <- struct{}{}
}
return &zwdaemonv1.SealResponse{}, s.sealErr
}
func (s *testService) ListFiles(
_ context.Context,
request *zwdaemonv1.ListFilesRequest,
) (*zwdaemonv1.ListFilesResponse, error) {
if s.listFilesRequests != nil {
s.listFilesRequests <- request
}
if s.listFilesResponse == nil {
return &zwdaemonv1.ListFilesResponse{}, s.listFilesErr
}
return s.listFilesResponse, s.listFilesErr
}
func (s *testService) DeleteFile(
_ context.Context,
request *zwdaemonv1.DeleteFileRequest,
) (*zwdaemonv1.DeleteFileResponse, error) {
if s.deleteFileRequests != nil {
s.deleteFileRequests <- request
}
return &zwdaemonv1.DeleteFileResponse{}, s.deleteFileErr
}
func (s *testService) Compact(
context.Context,
*zwdaemonv1.CompactRequest,
) (*zwdaemonv1.CompactResponse, error) {
if s.compactCalls != nil {
s.compactCalls <- struct{}{}
}
return &zwdaemonv1.CompactResponse{}, s.compactErr
}
func (s *testService) MoveFile(
_ context.Context,
request *zwdaemonv1.MoveFileRequest,
) (*zwdaemonv1.MoveFileResponse, error) {
if s.moveFileRequests != nil {
s.moveFileRequests <- request
}
return &zwdaemonv1.MoveFileResponse{}, s.moveFileErr
}
func (s *testService) CopyFile(
_ context.Context,
request *zwdaemonv1.CopyFileRequest,
) (*zwdaemonv1.CopyFileResponse, error) {
if s.copyFileRequests != nil {
s.copyFileRequests <- request
}
return &zwdaemonv1.CopyFileResponse{}, s.copyFileErr
}
func (s *testService) GetFile(
request *zwdaemonv1.GetFileRequest,
stream grpc.ServerStreamingServer[zwdaemonv1.GetFileResponse],
) error {
if s.getFile == nil {
return nil
}
return s.getFile(request, stream)
}
func (s *testService) GetFileRange(
request *zwdaemonv1.GetFileRangeRequest,
stream grpc.ServerStreamingServer[zwdaemonv1.GetFileRangeResponse],
) error {
if s.getFileRange == nil {
return nil
}
return s.getFileRange(request, stream)
}
func (s *testService) AddFile(
stream grpc.ClientStreamingServer[zwdaemonv1.AddFileRequest, zwdaemonv1.AddFileResponse],
) error {
if s.addFile == nil {
return stream.SendAndClose(&zwdaemonv1.AddFileResponse{})
}
return s.addFile(stream)
}
func (s *testService) ReplaceFile(
stream grpc.ClientStreamingServer[zwdaemonv1.ReplaceFileRequest, zwdaemonv1.ReplaceFileResponse],
) error {
if s.replaceFile == nil {
return stream.SendAndClose(&zwdaemonv1.ReplaceFileResponse{})
}
return s.replaceFile(stream)
}
func startTestServer(t *testing.T, service *testService) string {
t.Helper()
socketPath := filepath.Join(t.TempDir(), "daemon.sock")
listener, err := net.Listen("unix", socketPath)
if err != nil {
t.Fatalf("listen on Unix socket: %v", err)
}
server := grpc.NewServer()
zwdaemonv1.RegisterZwDaemonServiceServer(server, service)
serveErr := make(chan error, 1)
go func() {
serveErr <- server.Serve(listener)
}()
t.Cleanup(func() {
server.Stop()
if err := <-serveErr; err != nil {
t.Errorf("serve gRPC test server: %v", err)
}
})
return socketPath
}
func connectTestClient(t *testing.T, socketPath string) (*Client, error) {
t.Helper()
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
return Connect(ctx, socketPath)
}
func startConnectedTestClient(t *testing.T, service *testService) *Client {
t.Helper()
service.response = &zwdaemonv1.GetServerInfoResponse{ApiMajor: supportedAPIMajor}
client, err := connectTestClient(t, startTestServer(t, service))
if err != nil {
t.Fatalf("Connect() error = %v", err)
}
t.Cleanup(func() { _ = client.Close() })
return client
}
func TestConnectAcceptsSupportedMajorAndMinorDifferences(t *testing.T) {
for _, minor := range []uint32{0, 99} {
t.Run(fmt.Sprintf("minor_%d", minor), func(t *testing.T) {
socketPath := startTestServer(t, &testService{
response: &zwdaemonv1.GetServerInfoResponse{
ApiMajor: supportedAPIMajor,
ApiMinor: minor,
DaemonVersion: "test-version",
},
})
client, err := connectTestClient(t, socketPath)
if err != nil {
t.Fatalf("Connect() error = %v", err)
}
t.Cleanup(func() { _ = client.Close() })
})
}
}
func TestConnectRejectsUnsupportedMajor(t *testing.T) {
socketPath := startTestServer(t, &testService{
response: &zwdaemonv1.GetServerInfoResponse{
ApiMajor: 2,
ApiMinor: 3,
},
})
client, err := connectTestClient(t, socketPath)
if client != nil {
_ = client.Close()
t.Fatal("Connect() returned a client for an unsupported API major")
}
var versionErr *IncompatibleAPIVersionError
if !errors.As(err, &versionErr) {
t.Fatalf("Connect() error = %T %v, want *IncompatibleAPIVersionError", err, err)
}
if versionErr.ServerMajor != 2 || versionErr.ServerMinor != 3 || versionErr.SupportedMajor != 1 {
t.Fatalf("version error = %#v", versionErr)
}
}
func TestServerInfoMapsResponse(t *testing.T) {
socketPath := startTestServer(t, &testService{
response: &zwdaemonv1.GetServerInfoResponse{
ApiMajor: 1,
ApiMinor: 7,
DaemonVersion: "1.2.3-test",
},
})
client, err := connectTestClient(t, socketPath)
if err != nil {
t.Fatalf("Connect() error = %v", err)
}
t.Cleanup(func() { _ = client.Close() })
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
info, err := client.ServerInfo(ctx)
if err != nil {
t.Fatalf("ServerInfo() error = %v", err)
}
want := ServerInfo{APIMajor: 1, APIMinor: 7, DaemonVersion: "1.2.3-test"}
if info != want {
t.Fatalf("ServerInfo() = %#v, want %#v", info, want)
}
}
func TestConnectPreservesStructuredRPCError(t *testing.T) {
serverStatus, err := status.New(codes.FailedPrecondition, "vault is sealed").WithDetails(
&zwdaemonv1.ZwErrorDetail{
Code: zwdaemonv1.ZwErrorCode_ZW_ERROR_CODE_VAULT_SEALED,
Retryable: true,
},
)
if err != nil {
t.Fatalf("create structured status: %v", err)
}
socketPath := startTestServer(t, &testService{err: serverStatus.Err()})
_, err = connectTestClient(t, socketPath)
rpcErr := requireRPCError(t, err)
if rpcErr.GRPCCode != codes.FailedPrecondition {
t.Errorf("GRPCCode = %v, want %v", rpcErr.GRPCCode, codes.FailedPrecondition)
}
if rpcErr.ZwCode == nil || *rpcErr.ZwCode != ErrorCodeVaultSealed {
t.Errorf("ZwCode = %v, want %v", rpcErr.ZwCode, ErrorCodeVaultSealed)
}
if rpcErr.Message != "vault is sealed" {
t.Errorf("Message = %q, want %q", rpcErr.Message, "vault is sealed")
}
if !rpcErr.Retryable {
t.Error("Retryable = false, want true")
}
if got := status.Code(err); got != codes.FailedPrecondition {
t.Errorf("status.Code(error) = %v, want %v", got, codes.FailedPrecondition)
}
}
func TestConnectPreservesUnknownZwErrorCode(t *testing.T) {
const unknownCode zwdaemonv1.ZwErrorCode = 999
serverStatus, err := status.New(codes.Internal, "future daemon error").WithDetails(
&zwdaemonv1.ZwErrorDetail{Code: unknownCode},
)
if err != nil {
t.Fatalf("create structured status: %v", err)
}
socketPath := startTestServer(t, &testService{err: serverStatus.Err()})
_, err = connectTestClient(t, socketPath)
rpcErr := requireRPCError(t, err)
if rpcErr.ZwCode == nil || *rpcErr.ZwCode != unknownCode {
t.Fatalf("ZwCode = %v, want numeric value %d", rpcErr.ZwCode, unknownCode)
}
}
func TestConnectHandlesStatusWithoutZwDetail(t *testing.T) {
socketPath := startTestServer(t, &testService{
err: status.Error(codes.Unavailable, "daemon unavailable"),
})
_, err := connectTestClient(t, socketPath)
rpcErr := requireRPCError(t, err)
if rpcErr.ZwCode != nil {
t.Errorf("ZwCode = %v, want nil", rpcErr.ZwCode)
}
if rpcErr.Retryable {
t.Error("Retryable = true without ZwErrorDetail")
}
}
func TestConnectIgnoresMalformedZwErrorDetail(t *testing.T) {
serverErr := statusWithRawDetail(
codes.Internal,
"malformed detail",
testZwErrorDetailTypeURL,
[]byte{0xff},
)
socketPath := startTestServer(t, &testService{err: serverErr})
_, err := connectTestClient(t, socketPath)
rpcErr := requireRPCError(t, err)
if rpcErr.ZwCode != nil {
t.Errorf("ZwCode = %v, want nil", rpcErr.ZwCode)
}
if rpcErr.GRPCCode != codes.Internal || rpcErr.Message != "malformed detail" {
t.Errorf("base status was not preserved: %#v", rpcErr)
}
}
func TestConnectRequiresExactZwErrorDetailTypeURL(t *testing.T) {
encoded, err := proto.Marshal(&zwdaemonv1.ZwErrorDetail{
Code: zwdaemonv1.ZwErrorCode_ZW_ERROR_CODE_VAULT_SEALED,
Retryable: true,
})
if err != nil {
t.Fatalf("marshal ZwErrorDetail: %v", err)
}
serverErr := statusWithRawDetail(
codes.FailedPrecondition,
"wrong type URL prefix",
"https://example.invalid/zw.daemon.v1.ZwErrorDetail",
encoded,
)
socketPath := startTestServer(t, &testService{err: serverErr})
_, err = connectTestClient(t, socketPath)
rpcErr := requireRPCError(t, err)
if rpcErr.ZwCode != nil {
t.Errorf("ZwCode = %v for non-canonical type URL, want nil", rpcErr.ZwCode)
}
if rpcErr.Retryable {
t.Error("Retryable = true for non-canonical type URL")
}
}
func TestCreateAndUnsealMapRequests(t *testing.T) {
service := &testService{
createRequests: make(chan *zwdaemonv1.CreateRequest, 1),
unsealRequests: make(chan *zwdaemonv1.UnsealRequest, 1),
}
client := startConnectedTestClient(t, service)
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
if err := client.Create(ctx, "/vaults/new.zw", "create-password"); err != nil {
t.Fatalf("Create() error = %v", err)
}
createRequest := <-service.createRequests
if createRequest.GetVaultPath() != "/vaults/new.zw" {
t.Errorf("Create vault path = %q", createRequest.GetVaultPath())
}
if createRequest.GetPassword() != "create-password" {
t.Errorf("Create password = %q", createRequest.GetPassword())
}
if err := client.Unseal(ctx, "/vaults/existing.zw", "unseal-password"); err != nil {
t.Fatalf("Unseal() error = %v", err)
}
unsealRequest := <-service.unsealRequests
if unsealRequest.GetVaultPath() != "/vaults/existing.zw" {
t.Errorf("Unseal vault path = %q", unsealRequest.GetVaultPath())
}
if unsealRequest.GetPassword() != "unseal-password" {
t.Errorf("Unseal password = %q", unsealRequest.GetPassword())
}
}
func TestSealAndCompactCallDaemon(t *testing.T) {
service := &testService{
sealCalls: make(chan struct{}, 1),
compactCalls: make(chan struct{}, 1),
}
client := startConnectedTestClient(t, service)
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
if err := client.Seal(ctx); err != nil {
t.Fatalf("Seal() error = %v", err)
}
<-service.sealCalls
if err := client.Compact(ctx); err != nil {
t.Fatalf("Compact() error = %v", err)
}
<-service.compactCalls
}
func TestListFilesMapsFilterAndPreservesResponseOrder(t *testing.T) {
service := &testService{
listFilesRequests: make(chan *zwdaemonv1.ListFilesRequest, 1),
listFilesResponse: &zwdaemonv1.ListFilesResponse{
Files: []*zwdaemonv1.ListedFile{
{LogicalPath: "logs/second.txt", Size: 22},
{LogicalPath: "logs/first.txt", Size: 11},
},
},
}
client := startConnectedTestClient(t, service)
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
files, err := client.ListFiles(ctx, ListFilter{Prefix: "logs/", Contains: "txt"})
if err != nil {
t.Fatalf("ListFiles() error = %v", err)
}
request := <-service.listFilesRequests
if request.Prefix == nil || request.GetPrefix() != "logs/" {
t.Errorf("ListFiles prefix = %v", request.Prefix)
}
if request.Contains == nil || request.GetContains() != "txt" {
t.Errorf("ListFiles contains = %v", request.Contains)
}
want := []ListedFile{
{LogicalPath: "logs/second.txt", Size: 22},
{LogicalPath: "logs/first.txt", Size: 11},
}
if len(files) != len(want) {
t.Fatalf("ListFiles() returned %d files, want %d", len(files), len(want))
}
for index := range want {
if files[index] != want[index] {
t.Errorf("ListFiles()[%d] = %#v, want %#v", index, files[index], want[index])
}
}
}
func TestListPathsOmitsEmptyFiltersAndMapsPaths(t *testing.T) {
service := &testService{
listFilesRequests: make(chan *zwdaemonv1.ListFilesRequest, 1),
listFilesResponse: &zwdaemonv1.ListFilesResponse{
Files: []*zwdaemonv1.ListedFile{
{LogicalPath: "first", Size: 10},
{LogicalPath: "second", Size: 20},
},
},
}
client := startConnectedTestClient(t, service)
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
paths, err := client.ListPaths(ctx)
if err != nil {
t.Fatalf("ListPaths() error = %v", err)
}
request := <-service.listFilesRequests
if request.Prefix != nil || request.Contains != nil {
t.Errorf("ListPaths request filters = prefix %v, contains %v; want both absent", request.Prefix, request.Contains)
}
want := []string{"first", "second"}
if len(paths) != len(want) {
t.Fatalf("ListPaths() returned %d paths, want %d", len(paths), len(want))
}
for index := range want {
if paths[index] != want[index] {
t.Errorf("ListPaths()[%d] = %q, want %q", index, paths[index], want[index])
}
}
}
func TestDeleteFileMapsRequestAndStructuredError(t *testing.T) {
serverStatus, err := status.New(codes.NotFound, "entry not found").WithDetails(
&zwdaemonv1.ZwErrorDetail{
Code: zwdaemonv1.ZwErrorCode_ZW_ERROR_CODE_ENTRY_NOT_FOUND,
},
)
if err != nil {
t.Fatalf("create structured status: %v", err)
}
service := &testService{
deleteFileRequests: make(chan *zwdaemonv1.DeleteFileRequest, 1),
deleteFileErr: serverStatus.Err(),
}
client := startConnectedTestClient(t, service)
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
err = client.DeleteFile(ctx, "missing.txt")
request := <-service.deleteFileRequests
if request.GetLogicalPath() != "missing.txt" {
t.Errorf("DeleteFile logical path = %q", request.GetLogicalPath())
}
rpcErr := requireRPCError(t, err)
if rpcErr.GRPCCode != codes.NotFound {
t.Errorf("GRPCCode = %v, want %v", rpcErr.GRPCCode, codes.NotFound)
}
if rpcErr.ZwCode == nil || *rpcErr.ZwCode != ErrorCodeEntryNotFound {
t.Errorf("ZwCode = %v, want %v", rpcErr.ZwCode, ErrorCodeEntryNotFound)
}
}
func TestMoveAndCopyFileMapRequests(t *testing.T) {
service := &testService{
moveFileRequests: make(chan *zwdaemonv1.MoveFileRequest, 1),
copyFileRequests: make(chan *zwdaemonv1.CopyFileRequest, 1),
}
client := startConnectedTestClient(t, service)
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
if err := client.MoveFile(ctx, "source.txt", "renamed.txt"); err != nil {
t.Fatalf("MoveFile() error = %v", err)
}
moveRequest := <-service.moveFileRequests
if moveRequest.GetSourcePath() != "source.txt" || moveRequest.GetDestinationPath() != "renamed.txt" {
t.Errorf("MoveFile request = %#v", moveRequest)
}
if err := client.CopyFile(ctx, "source.txt", "copy.txt"); err != nil {
t.Fatalf("CopyFile() error = %v", err)
}
copyRequest := <-service.copyFileRequests
if copyRequest.GetSourcePath() != "source.txt" || copyRequest.GetDestinationPath() != "copy.txt" {
t.Errorf("CopyFile request = %#v", copyRequest)
}
}
func TestConvertRPCErrorAcceptsNil(t *testing.T) {
if err := convertRPCError(nil); err != nil {
t.Fatalf("convertRPCError(nil) = %v", err)
}
}
func TestGetFileReadsChunksAcrossBuffersAndSkipsEmptyChunks(t *testing.T) {
requests := make(chan *zwdaemonv1.GetFileRequest, 1)
service := &testService{
getFile: func(
request *zwdaemonv1.GetFileRequest,
stream grpc.ServerStreamingServer[zwdaemonv1.GetFileResponse],
) error {
requests <- request
for _, chunk := range [][]byte{[]byte("abc"), nil, []byte("defgh")} {
if err := stream.Send(&zwdaemonv1.GetFileResponse{Data: chunk}); err != nil {
return err
}
}
return nil
},
}
client := startConnectedTestClient(t, service)
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
reader, err := client.GetFile(ctx, "folder/file.bin")
if err != nil {
t.Fatalf("GetFile() error = %v", err)
}
t.Cleanup(func() { _ = reader.Close() })
request := <-requests
if request.GetLogicalPath() != "folder/file.bin" {
t.Errorf("GetFile logical path = %q", request.GetLogicalPath())
}
buffer := make([]byte, 2)
var contents []byte
for {
n, err := reader.Read(buffer)
contents = append(contents, buffer[:n]...)
if err == io.EOF {
break
}
if err != nil {
t.Fatalf("Read() error = %v", err)
}
if n == 0 {
t.Fatal("Read() returned no bytes and no error")
}
}
if got, want := string(contents), "abcdefgh"; got != want {
t.Fatalf("download contents = %q, want %q", got, want)
}
}
func TestGetFileRangeMapsRequestAndReadsResponse(t *testing.T) {
requests := make(chan *zwdaemonv1.GetFileRangeRequest, 1)
service := &testService{
getFileRange: func(
request *zwdaemonv1.GetFileRangeRequest,
stream grpc.ServerStreamingServer[zwdaemonv1.GetFileRangeResponse],
) error {
requests <- request
return stream.Send(&zwdaemonv1.GetFileRangeResponse{Data: []byte("range")})
},
}
client := startConnectedTestClient(t, service)
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
reader, err := client.GetFileRange(ctx, "large.bin", 10, 15)
if err != nil {
t.Fatalf("GetFileRange() error = %v", err)
}
t.Cleanup(func() { _ = reader.Close() })
contents, err := io.ReadAll(reader)
if err != nil {
t.Fatalf("read range: %v", err)
}
if got, want := string(contents), "range"; got != want {
t.Fatalf("range contents = %q, want %q", got, want)
}
request := <-requests
if request.GetLogicalPath() != "large.bin" || request.GetStart() != 10 || request.GetExclusiveEnd() != 15 {
t.Fatalf("GetFileRange request = %#v", request)
}
}
func TestGetFileReturnsTerminalStructuredErrorAfterReceivedData(t *testing.T) {
serverStatus, err := status.New(codes.DataLoss, "download failed").WithDetails(
&zwdaemonv1.ZwErrorDetail{
Code: zwdaemonv1.ZwErrorCode_ZW_ERROR_CODE_INTERNAL,
Retryable: true,
},
)
if err != nil {
t.Fatalf("create structured status: %v", err)
}
service := &testService{
getFile: func(
_ *zwdaemonv1.GetFileRequest,
stream grpc.ServerStreamingServer[zwdaemonv1.GetFileResponse],
) error {
if err := stream.Send(&zwdaemonv1.GetFileResponse{Data: []byte("partial")}); err != nil {
return err
}
return serverStatus.Err()
},
}
client := startConnectedTestClient(t, service)
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
reader, err := client.GetFile(ctx, "damaged.bin")
if err != nil {
t.Fatalf("GetFile() error = %v", err)
}
t.Cleanup(func() { _ = reader.Close() })
buffer := make([]byte, 32)
n, err := reader.Read(buffer)
if err != nil {
t.Fatalf("first Read() error = %v", err)
}
if got, want := string(buffer[:n]), "partial"; got != want {
t.Fatalf("first Read() = %q, want %q", got, want)
}
n, err = reader.Read(buffer)
if n != 0 {
t.Errorf("terminal Read() returned %d bytes", n)
}
rpcErr := requireRPCError(t, err)
if rpcErr.GRPCCode != codes.DataLoss {
t.Errorf("GRPCCode = %v, want %v", rpcErr.GRPCCode, codes.DataLoss)
}
if rpcErr.ZwCode == nil || *rpcErr.ZwCode != ErrorCodeInternal {
t.Errorf("ZwCode = %v, want %v", rpcErr.ZwCode, ErrorCodeInternal)
}
if rpcErr.Message != "download failed" || !rpcErr.Retryable {
t.Errorf("terminal RPC error = %#v", rpcErr)
}
infoCtx, infoCancel := context.WithTimeout(context.Background(), 5*time.Second)
defer infoCancel()
if _, err := client.ServerInfo(infoCtx); err != nil {
t.Fatalf("ServerInfo() after failed download = %v", err)
}
}
func TestClosingDownloadCancelsStreamAndKeepsClientReusable(t *testing.T) {
started := make(chan struct{}, 1)
canceled := make(chan struct{}, 1)
service := &testService{
getFile: func(
_ *zwdaemonv1.GetFileRequest,
stream grpc.ServerStreamingServer[zwdaemonv1.GetFileResponse],
) error {
started <- struct{}{}
<-stream.Context().Done()
canceled <- struct{}{}
return stream.Context().Err()
},
}
client := startConnectedTestClient(t, service)
reader, err := client.GetFile(context.Background(), "abandoned.bin")
if err != nil {
t.Fatalf("GetFile() error = %v", err)
}
select {
case <-started:
case <-time.After(5 * time.Second):
t.Fatal("server did not start download")
}
if err := reader.Close(); err != nil {
t.Fatalf("Close() error = %v", err)
}
if err := reader.Close(); err != nil {
t.Fatalf("second Close() error = %v", err)
}
if _, err := reader.Read(make([]byte, 1)); !errors.Is(err, fs.ErrClosed) {
t.Fatalf("Read() after Close() error = %v, want fs.ErrClosed", err)
}
select {
case <-canceled:
case <-time.After(5 * time.Second):
t.Fatal("closing reader did not cancel server stream")
}
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
if _, err := client.ServerInfo(ctx); err != nil {
t.Fatalf("ServerInfo() after closing download = %v", err)
}
}
func TestAddFileSendsHeaderBeforeEmptyAndMultiChunkBodies(t *testing.T) {
tests := []struct {
name string
contents []byte
chunkSizes []int
}{
{name: "empty"},
{
name: "multiple_chunks",
contents: bytes.Repeat([]byte("x"), 512*1024+137),
chunkSizes: []int{512 * 1024, 137},
},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
received := make(chan receivedUpload, 1)
service := &testService{
addFile: func(
stream grpc.ClientStreamingServer[zwdaemonv1.AddFileRequest, zwdaemonv1.AddFileResponse],
) error {
upload, err := receiveUpload(stream)
if err != nil {
return err
}
received <- upload
return stream.SendAndClose(&zwdaemonv1.AddFileResponse{})
},
}
client := startConnectedTestClient(t, service)
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
err := client.AddFile(ctx, "upload.bin", bytes.NewReader(test.contents), uint64(len(test.contents)))
if err != nil {
t.Fatalf("AddFile() error = %v", err)
}
upload := <-received
if upload.header.GetLogicalPath() != "upload.bin" {
t.Errorf("logical path = %q", upload.header.GetLogicalPath())
}
if upload.header.GetSourceLen() != uint64(len(test.contents)) {
t.Errorf("source length = %d", upload.header.GetSourceLen())
}
if !bytes.Equal(upload.contents, test.contents) {
t.Errorf("uploaded contents differ")
}
if len(upload.chunkSizes) != len(test.chunkSizes) {
t.Fatalf("chunk sizes = %v, want %v", upload.chunkSizes, test.chunkSizes)
}
for index := range test.chunkSizes {
if upload.chunkSizes[index] != test.chunkSizes[index] {
t.Errorf("chunk sizes = %v, want %v", upload.chunkSizes, test.chunkSizes)
}
}
})
}
}
func TestAddFileProcessesDataReturnedWithEOF(t *testing.T) {
received := make(chan receivedUpload, 1)
service := &testService{
addFile: func(
stream grpc.ClientStreamingServer[zwdaemonv1.AddFileRequest, zwdaemonv1.AddFileResponse],
) error {
upload, err := receiveUpload(stream)
if err != nil {
return err
}
received <- upload
return stream.SendAndClose(&zwdaemonv1.AddFileResponse{})
},
}
client := startConnectedTestClient(t, service)
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
err := client.AddFile(ctx, "combined.bin", &dataAndEOFReader{data: []byte("contents")}, 8)
if err != nil {
t.Fatalf("AddFile() error = %v", err)
}
upload := <-received
if got, want := string(upload.contents), "contents"; got != want {
t.Fatalf("uploaded contents = %q, want %q", got, want)
}
}
func TestReplaceFileSendsHeaderAndChunks(t *testing.T) {
received := make(chan receivedReplaceUpload, 1)
service := &testService{
replaceFile: func(
stream grpc.ClientStreamingServer[zwdaemonv1.ReplaceFileRequest, zwdaemonv1.ReplaceFileResponse],
) error {
upload, err := receiveReplaceUpload(stream)
if err != nil {
return err
}
received <- upload
return stream.SendAndClose(&zwdaemonv1.ReplaceFileResponse{})
},
}
client := startConnectedTestClient(t, service)
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
contents := []byte("replacement contents")
if err := client.ReplaceFile(ctx, "existing.bin", bytes.NewReader(contents), uint64(len(contents))); err != nil {
t.Fatalf("ReplaceFile() error = %v", err)
}
upload := <-received
if upload.header.GetLogicalPath() != "existing.bin" {
t.Errorf("logical path = %q", upload.header.GetLogicalPath())
}
if upload.header.GetSourceLen() != uint64(len(contents)) {
t.Errorf("source length = %d", upload.header.GetSourceLen())
}
if !bytes.Equal(upload.contents, contents) {
t.Errorf("replacement contents differ")
}
}
func TestAddFileReturnsLengthRejectionFromCloseAndRecv(t *testing.T) {
serverStatus, err := status.New(codes.InvalidArgument, "declared length does not match source").WithDetails(
&zwdaemonv1.ZwErrorDetail{
Code: zwdaemonv1.ZwErrorCode_ZW_ERROR_CODE_INVALID_REQUEST,
},
)
if err != nil {
t.Fatalf("create structured status: %v", err)
}
service := &testService{
addFile: func(
stream grpc.ClientStreamingServer[zwdaemonv1.AddFileRequest, zwdaemonv1.AddFileResponse],
) error {
if _, err := receiveUpload(stream); err != nil {
return err
}
return serverStatus.Err()
},
}
client := startConnectedTestClient(t, service)
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
err = client.AddFile(ctx, "short.bin", bytes.NewReader([]byte("abc")), 4)
rpcErr := requireRPCError(t, err)
if rpcErr.GRPCCode != codes.InvalidArgument {
t.Errorf("GRPCCode = %v, want %v", rpcErr.GRPCCode, codes.InvalidArgument)
}
if rpcErr.ZwCode == nil || *rpcErr.ZwCode != ErrorCodeInvalidRequest {
t.Errorf("ZwCode = %v, want %v", rpcErr.ZwCode, ErrorCodeInvalidRequest)
}
}
func TestUploadReadFailureCancelsRPCAndConnectionAndPathRemainReusable(t *testing.T) {
injectedErr := errors.New("injected source error")
firstStreamErr := make(chan error, 1)
headers := make(chan *zwdaemonv1.AddFileHeader, 2)
var calls atomic.Int32
service := &testService{
addFile: func(
stream grpc.ClientStreamingServer[zwdaemonv1.AddFileRequest, zwdaemonv1.AddFileResponse],
) error {
call := calls.Add(1)
if call == 1 {
first, err := stream.Recv()
if err != nil {
firstStreamErr <- err
return err
}
headers <- first.GetHeader()
for {
_, err := stream.Recv()
if err != nil {
firstStreamErr <- err
return err
}
}
}
upload, err := receiveUpload(stream)
if err != nil {
return err
}
headers <- upload.header
return stream.SendAndClose(&zwdaemonv1.AddFileResponse{})
},
}
client := startConnectedTestClient(t, service)
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
err := client.AddFile(ctx, "reusable.bin", &dataAndErrorReader{
data: []byte("partial"),
err: injectedErr,
}, 7)
var readErr *UploadReadError
if !errors.As(err, &readErr) {
t.Fatalf("AddFile() error = %T %v, want *UploadReadError", err, err)
}
if !errors.Is(err, injectedErr) {
t.Fatalf("AddFile() error does not wrap source error: %v", err)
}
select {
case err := <-firstStreamErr:
if errors.Is(err, io.EOF) {
t.Fatal("reader failure cleanly closed the upload stream")
}
if !errors.Is(err, context.Canceled) && status.Code(err) != codes.Canceled {
t.Fatalf("server stream error = %v, want cancellation", err)
}
case <-time.After(5 * time.Second):
t.Fatal("server did not observe upload cancellation")
}
infoCtx, infoCancel := context.WithTimeout(context.Background(), 5*time.Second)
defer infoCancel()
if _, err := client.ServerInfo(infoCtx); err != nil {
t.Fatalf("ServerInfo() after reader failure = %v", err)
}
if err := client.AddFile(ctx, "reusable.bin", bytes.NewReader([]byte("replacement")), 11); err != nil {
t.Fatalf("second AddFile() on same path = %v", err)
}
for index := 0; index < 2; index++ {
header := <-headers
if header == nil || header.GetLogicalPath() != "reusable.bin" {
t.Fatalf("upload header = %#v", header)
}
}
}
type receivedUpload struct {
header *zwdaemonv1.AddFileHeader
contents []byte
chunkSizes []int
}
func receiveUpload(
stream grpc.ClientStreamingServer[zwdaemonv1.AddFileRequest, zwdaemonv1.AddFileResponse],
) (receivedUpload, error) {
first, err := stream.Recv()
if err != nil {
return receivedUpload{}, err
}
header := first.GetHeader()
if header == nil {
return receivedUpload{}, status.Error(codes.InvalidArgument, "first upload message is not a header")
}
upload := receivedUpload{header: header}
for {
request, err := stream.Recv()
if err == io.EOF {
return upload, nil
}
if err != nil {
return upload, err
}
payload, ok := request.GetPayload().(*zwdaemonv1.AddFileRequest_Chunk)
if !ok {
return upload, status.Error(codes.InvalidArgument, "upload message after header is not a chunk")
}
upload.contents = append(upload.contents, payload.Chunk...)
upload.chunkSizes = append(upload.chunkSizes, len(payload.Chunk))
}
}
type receivedReplaceUpload struct {
header *zwdaemonv1.ReplaceFileHeader
contents []byte
}
func receiveReplaceUpload(
stream grpc.ClientStreamingServer[zwdaemonv1.ReplaceFileRequest, zwdaemonv1.ReplaceFileResponse],
) (receivedReplaceUpload, error) {
first, err := stream.Recv()
if err != nil {
return receivedReplaceUpload{}, err
}
header := first.GetHeader()
if header == nil {
return receivedReplaceUpload{}, status.Error(codes.InvalidArgument, "first replacement message is not a header")
}
upload := receivedReplaceUpload{header: header}
for {
request, err := stream.Recv()
if err == io.EOF {
return upload, nil
}
if err != nil {
return upload, err
}
payload, ok := request.GetPayload().(*zwdaemonv1.ReplaceFileRequest_Chunk)
if !ok {
return upload, status.Error(codes.InvalidArgument, "replacement message after header is not a chunk")
}
upload.contents = append(upload.contents, payload.Chunk...)
}
}
type dataAndEOFReader struct {
data []byte
done bool
}
func (r *dataAndEOFReader) Read(destination []byte) (int, error) {
if r.done {
return 0, io.EOF
}
r.done = true
return copy(destination, r.data), io.EOF
}
type dataAndErrorReader struct {
data []byte
err error
done bool
}
func (r *dataAndErrorReader) Read(destination []byte) (int, error) {
if r.done {
return 0, r.err
}
r.done = true
return copy(destination, r.data), r.err
}
func requireRPCError(t *testing.T, err error) *RPCError {
t.Helper()
var rpcErr *RPCError
if !errors.As(err, &rpcErr) {
t.Fatalf("error = %T %v, want *RPCError", err, err)
}
return rpcErr
}
func statusWithRawDetail(code codes.Code, message, typeURL string, value []byte) error {
statusProto := status.New(code, message).Proto()
statusProto.Details = append(statusProto.Details, &anypb.Any{
TypeUrl: typeURL,
Value: value,
})
return status.FromProto(statusProto).Err()
}