diff --git a/src/client/client_send.c b/src/client/client_send.c index bf19384..f32e569 100644 --- a/src/client/client_send.c +++ b/src/client/client_send.c @@ -16,6 +16,8 @@ #include "transport_ssh.h" #include "transport_tls.h" #include "utils.h" +#include +#include #include #include #include @@ -105,26 +107,83 @@ static bool finalize_transfer(Client* client) { receive_status(client->file_descriptor, &status) && status == STATUS_OK; } -/* Remove only regular source files after the receiver confirms the whole transfer. */ +typedef struct { + char* path; + dev_t device; + ino_t inode; +} SourceFile; + +static void source_file_destroy(void* item) { + SourceFile* source = item; + if (source) { + free(source->path); + free(source); + } +} + +/* Remove only the same regular source file that was sent. */ static void remove_transferred_sources(const Config* config, ArrayList* paths) { if (!config->remove_source_files || !paths) return; for (int i = 0; i < paths->size; i++) { - const char* path = paths->items[i]; - struct stat st; - if (lstat(path, &st) != 0 || !S_ISREG(st.st_mode)) + SourceFile* source = paths->items[i]; + const char* slash = strrchr(source->path, '/'); + const char* leaf = slash ? slash + 1 : source->path; + char parent[PATH_MAX]; + if (slash) { + size_t parent_length = (size_t)(slash - source->path); + if (parent_length == 0) + parent_length = 1; + if (parent_length >= sizeof(parent)) + continue; + memcpy(parent, source->path, parent_length); + parent[parent_length] = '\0'; + } else { + (void)snprintf(parent, sizeof(parent), "."); + } + + int dirfd = open(parent, O_RDONLY | O_DIRECTORY | O_CLOEXEC); + if (dirfd < 0) continue; - if (unlink(path) != 0) - log_message(LOG_LEVEL_WARNING, "Could not remove source file %s", path); + struct stat st; + if (fstatat(dirfd, leaf, &st, AT_SYMLINK_NOFOLLOW) != 0 || !S_ISREG(st.st_mode) || + st.st_dev != source->device || st.st_ino != source->inode) { + close(dirfd); + continue; + } + if (unlinkat(dirfd, leaf, 0) != 0) + log_message(LOG_LEVEL_WARNING, "Could not remove source file %s", source->path); + close(dirfd); } } +static SourceFile* source_file_create(const File* file) { + if (!file || !file->path) + return NULL; + struct stat st; + if (lstat(file->path, &st) != 0 || !S_ISREG(st.st_mode)) + return NULL; + SourceFile* source = malloc(sizeof(*source)); + if (!source) + return NULL; + source->path = str_dup(file->path); + source->device = st.st_dev; + source->inode = st.st_ino; + if (!source->path) { + source_file_destroy(source); + return NULL; + } + return source; +} + static bool remember_source_file(ArrayList* paths, const File* file) { if (!paths || !file || !file->path) return true; - char* path = str_dup(file->path); - if (!path || !array_list_add(paths, path)) { - free(path); + SourceFile* source = source_file_create(file); + if (!source) + return true; + if (!array_list_add(paths, source)) { + source_file_destroy(source); return false; } return true; @@ -372,6 +431,12 @@ static int send_single_file(Client* client, File* file, Config* config, bool use static int send_chunk_with_removal(Client* client, Chunk* chunk, Config* config, ArrayList* remove_sources) { if (config->use_chunk_serialization) { + if (remove_sources) { + for (int i = 0; i < chunk->element_count; i++) { + if (!remember_source_file(remove_sources, chunk->items[i])) + return -1; + } + } if (!send_status(client->file_descriptor, STATUS_CHUNK)) return -1; Data* data; @@ -387,12 +452,6 @@ static int send_chunk_with_removal(Client* client, Chunk* chunk, Config* config, return -1; } data_destroy(data); - if (remove_sources) { - for (int i = 0; i < chunk->element_count; i++) { - if (!remember_source_file(remove_sources, chunk->items[i])) - return -1; - } - } return 0; } @@ -403,13 +462,20 @@ static int send_chunk_with_removal(Client* client, Chunk* chunk, Config* config, bool stream = f->data->data == NULL && f->data->size > 0; bool use_sendfile = (config->use_sendfile && !config->use_compression) || (stream && !config->use_compression); + SourceFile* source = remove_sources ? source_file_create(f) : NULL; int rc = send_single_file(client, f, config, config->use_incremental, use_sendfile); - if (rc == 1) + if (rc == 1) { + source_file_destroy(source); continue; - if (rc < 0) + } + if (rc < 0) { + source_file_destroy(source); return -1; - if (rc == 0 && !remember_source_file(remove_sources, f)) + } + if (source && !array_list_add(remove_sources, source)) { + source_file_destroy(source); return -1; + } } return 0; } @@ -652,7 +718,7 @@ int send_files(Config* config) { scanner = directory_scanner_create_with_options(config->send_directory, &scanner_options); manifest = create_transfer_manifest(config); if (config->remove_source_files) - remove_sources = array_list_create(free); + remove_sources = array_list_create(source_file_destroy); if (!scanner || (config->use_delete && !manifest) || (config->remove_source_files && !remove_sources)) goto send_fail; @@ -782,7 +848,7 @@ int send_files_multithreaded(Config** config_ptr) { if (config->use_delete) context->manifest = array_list_create(free); if (config->remove_source_files) - context->remove_source_files = array_list_create(free); + context->remove_source_files = array_list_create(source_file_destroy); if ((config->use_delete && !context->manifest) || (config->remove_source_files && !context->remove_source_files)) { pipeline_context_sender_destroy(context); diff --git a/tests/integration/test_features.py b/tests/integration/test_features.py index ee3690c..45cd0d3 100644 --- a/tests/integration/test_features.py +++ b/tests/integration/test_features.py @@ -57,6 +57,20 @@ class TestRemoveSourceFiles: assert os.path.isdir(os.path.join(source, "directory")) assert os.path.islink(os.path.join(source, "link.txt")) + def test_single_threaded_removes_transferred_file(self, shared_server): + source = os.path.join(TEST_DATA_DIR, "remove_single_source") + dest = os.path.join(TEST_DATA_DIR, "remove_single_dest") + clean_dir(source) + clean_dir(dest) + source_file = os.path.join(source, "file.txt") + with open(source_file, "wb") as f: + f.write(b"single threaded") + + result, _ = run_client(source, dest, flags=["--remove-source-files"], + port=shared_server.port) + assert result.returncode == 0 + assert not os.path.exists(source_file) + def test_dry_run_preserves_source_files(self): source = os.path.join(TEST_DATA_DIR, "remove_dry_source") dest = os.path.join(TEST_DATA_DIR, "remove_dry_dest") @@ -83,6 +97,23 @@ class TestRemoveSourceFiles: assert result.returncode != 0 assert os.path.isfile(source_file) + def test_incremental_skip_preserves_source_file(self, shared_server): + source = os.path.join(TEST_DATA_DIR, "remove_skipped_source") + dest = os.path.join(TEST_DATA_DIR, "remove_skipped_dest") + clean_dir(source) + clean_dir(dest) + source_file = os.path.join(source, "file.txt") + with open(source_file, "wb") as f: + f.write(b"keep after skip") + + result, _ = run_client(source, dest, port=shared_server.port) + assert result.returncode == 0 + result, _ = run_client(source, dest, + flags=["--remove-source-files", "--incremental"], + port=shared_server.port) + assert result.returncode == 0 + assert os.path.isfile(source_file) + class TestArchiveMode: def test_archive_mode(self, shared_server):