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 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 } 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) 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 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 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 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 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() }