1154 lines
33 KiB
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()
|
|
}
|