diff --git a/src/client/client_cli.c b/src/client/client_cli.c index 9ff7bf1..c3232b9 100644 --- a/src/client/client_cli.c +++ b/src/client/client_cli.c @@ -53,7 +53,7 @@ int main(int argc, char *argv[]) { save_to_disk = true; } - Config *config = config_create(str_dup("1.0.0"), NULL, NULL, + Config *config = config_create(str_dup(PROTOCOL_VERSION), NULL, NULL, save_to_disk, false, false, false, false, 5, false, 0); int positional_args[2]; diff --git a/src/server/server.c b/src/server/server.c index 577b430..25ac13f 100644 --- a/src/server/server.c +++ b/src/server/server.c @@ -26,6 +26,11 @@ int receive_files(Config *config, int file_descriptor) { if (config->use_compression) { data_to_process = data_decompress(chunk_data); data_destroy(chunk_data); + if (data_to_process == NULL) { + log_message(LOG_LEVEL_ERROR, "Failed to decompress chunk"); + send_status(file_descriptor, STATUS_ERROR); + return -1; + } } Chunk *chunk = chunk_deserialize(data_to_process, config->use_metadata); data_destroy(data_to_process); @@ -95,6 +100,16 @@ void handler(int file_descriptor) { close(file_descriptor); } +static Server *g_server = NULL; + +static void cleanup(int sig) { + (void)sig; + if (g_server) { + server_delete(&g_server); + } + _exit(0); +} + int main(int argc, char *argv[]) { signal(SIGPIPE, SIG_IGN); for (int i = 1; i < argc; i++) { @@ -106,8 +121,9 @@ int main(int argc, char *argv[]) { set_log_level(LOG_LEVEL_DEBUG); } } - Server *server = server_create(8080); - server_listen(server, handler); - server_delete(&server); + signal(SIGINT, cleanup); + signal(SIGTERM, cleanup); + g_server = server_create(8080); + server_listen(g_server, handler); return 0; } diff --git a/src/shared/compression.c b/src/shared/compression.c index cc605e8..8ff5de2 100644 --- a/src/shared/compression.c +++ b/src/shared/compression.c @@ -4,15 +4,22 @@ #include "stdlib.h" #include "zstd.h" +#define INITIAL_DECOMPRESS_BUF_SIZE (1024 * 1024) + Data *data_compress(Data *data_to_compress, int compression_level) { log_message(LOG_LEVEL_DEBUG, "Starting to compress data"); size_t dst_size = ZSTD_compressBound(data_to_compress->size); Data *compressed_data = data_create_empty(dst_size); + if (!compressed_data) { + log_message(LOG_LEVEL_ERROR, "Failed to allocate compression buffer"); + return NULL; + } ZSTD_CCtx *cctx = ZSTD_createCCtx(); if (!cctx) { log_message(LOG_LEVEL_ERROR, "Failed to create ZSTD compression context"); - exit(EXIT_FAILURE); + data_destroy(compressed_data); + return NULL; } ZSTD_inBuffer input = {data_to_compress->data, data_to_compress->size, 0}; @@ -24,7 +31,9 @@ Data *data_compress(Data *data_to_compress, int compression_level) { if (ZSTD_isError(ret)) { log_message(LOG_LEVEL_ERROR, "Compression failed: %s", ZSTD_getErrorName(ret)); - exit(EXIT_FAILURE); + ZSTD_freeCCtx(cctx); + data_destroy(compressed_data); + return NULL; } } while (ret > 0); @@ -40,23 +49,26 @@ Data *data_decompress(Data *compressed_data) { log_message(LOG_LEVEL_DEBUG, "Start to decompress data"); unsigned long long dst_size = ZSTD_getFrameContentSize( compressed_data->data, compressed_data->size); - if (ZSTD_isError(dst_size)) { - log_message(LOG_LEVEL_ERROR, "Failed to get decompressed size: %s", - ZSTD_getErrorName(dst_size)); - exit(EXIT_FAILURE); - } - - Data *uncompressed_data = data_create_empty((size_t)dst_size); ZSTD_DCtx *dctx = ZSTD_createDCtx(); if (!dctx) { log_message(LOG_LEVEL_ERROR, "Failed to create ZSTD decompression context"); - exit(EXIT_FAILURE); + return NULL; + } + + size_t buf_size = (!ZSTD_isError(dst_size) && dst_size > 0) + ? (size_t)dst_size + : INITIAL_DECOMPRESS_BUF_SIZE; + Data *uncompressed_data = data_create_empty(buf_size); + if (!uncompressed_data) { + log_message(LOG_LEVEL_ERROR, "Failed to allocate decompression buffer"); + ZSTD_freeDCtx(dctx); + return NULL; } ZSTD_inBuffer input = {compressed_data->data, compressed_data->size, 0}; - ZSTD_outBuffer output = {uncompressed_data->data, (size_t)dst_size, 0}; + ZSTD_outBuffer output = {uncompressed_data->data, buf_size, 0}; size_t ret; do { @@ -64,7 +76,22 @@ Data *data_decompress(Data *compressed_data) { if (ZSTD_isError(ret)) { log_message(LOG_LEVEL_ERROR, "Decompression failed: %s", ZSTD_getErrorName(ret)); - exit(EXIT_FAILURE); + ZSTD_freeDCtx(dctx); + data_destroy(uncompressed_data); + return NULL; + } + if (ret > 0 && output.pos == output.size) { + buf_size *= 2; + void *new_data = realloc(uncompressed_data->data, buf_size); + if (!new_data) { + log_message(LOG_LEVEL_ERROR, "Failed to grow decompression buffer"); + ZSTD_freeDCtx(dctx); + data_destroy(uncompressed_data); + return NULL; + } + uncompressed_data->data = new_data; + output.dst = new_data; + output.size = buf_size; } } while (ret > 0); diff --git a/src/shared/config.c b/src/shared/config.c index e5d8637..f4bbbfe 100644 --- a/src/shared/config.c +++ b/src/shared/config.c @@ -90,6 +90,14 @@ void config_send(int file_descriptor, Config *config) { Config *config_receive(int file_descriptor) { Config *config = (Config *)malloc(sizeof(Config)); config->version = receive_str(file_descriptor); + if (strcmp(config->version, PROTOCOL_VERSION) != 0) { + fprintf(stderr, "Protocol version mismatch: client=%s, server=%s\n", + config->version, PROTOCOL_VERSION); + free(config->version); + free(config); + send_status(file_descriptor, STATUS_ERROR); + exit(EXIT_FAILURE); + } config->send_directory = receive_str(file_descriptor); config->receive_root_directory = receive_str(file_descriptor); config->save_to_disk = receive_int(file_descriptor); diff --git a/src/shared/config.h b/src/shared/config.h index b1c2437..58f452e 100644 --- a/src/shared/config.h +++ b/src/shared/config.h @@ -30,6 +30,7 @@ typedef struct Config { int exclude_count; } Config; +#define PROTOCOL_VERSION "1.0.0" #define DEFAULT_CHUNK_SIZE (10 * 1024 * 1024) Config *config_create(char *version, char *send_directory, diff --git a/src/shared/file.c b/src/shared/file.c index a04b177..6b4c87f 100644 --- a/src/shared/file.c +++ b/src/shared/file.c @@ -10,6 +10,7 @@ #include #include "compression.h" +#include "log.h" #include "config.h" #include "data.h" #include "file.h" @@ -95,6 +96,10 @@ void file_send_single_calls(File *file, int file_descriptor, bool use_metadata, if (compression_level > 0) { Data *compressed_data = data_compress(file->data, compression_level); data_destroy(file->data); + if (compressed_data == NULL) { + log_message(LOG_LEVEL_ERROR, "Compression failed in file_send_single_calls"); + exit(EXIT_FAILURE); + } file->data = compressed_data; } send_str(file_descriptor, file->path); diff --git a/src/shared/multiprocessing.c b/src/shared/multiprocessing.c index bcf0812..912be13 100644 --- a/src/shared/multiprocessing.c +++ b/src/shared/multiprocessing.c @@ -85,6 +85,10 @@ static void receive_chunk_enqueue(int file_descriptor, if (context->config->use_compression) { data_to_process = data_decompress(chunk_data); data_destroy(chunk_data); + if (data_to_process == NULL) { + log_message(LOG_LEVEL_ERROR, "Failed to decompress chunk, skipping"); + return; + } } Chunk *chunk = chunk_deserialize(data_to_process, context->config->use_metadata); data_destroy(data_to_process); diff --git a/src/shared/transport_tcp.c b/src/shared/transport_tcp.c index 1349a9a..9e8dd92 100644 --- a/src/shared/transport_tcp.c +++ b/src/shared/transport_tcp.c @@ -1,6 +1,7 @@ #include "transport_tcp.h" #include "log.h" #include +#include #include #include #include @@ -48,6 +49,7 @@ Server *server_create(int port) { void server_delete(Server **server) { if (server == NULL || *server == NULL) return; + close((*server)->file_descriptor); free(*server); *server = NULL; } @@ -55,22 +57,33 @@ void server_delete(Server **server) { void server_listen(Server *server, void (*handler)(int file_descriptor)) { log_message(LOG_LEVEL_INFO, "Start Listening on Port: %d", server->address.sin_port); - if (listen(server->file_descriptor, 3) < 0) { + if (listen(server->file_descriptor, SOMAXCONN) < 0) { perror("Could not listen on port!"); exit(EXIT_FAILURE); } - int file_descriptor = - accept(server->file_descriptor, (struct sockaddr *)&server->address, - &server->address_length); - if (file_descriptor < 0) { - perror("Could not accept the connection"); - exit(EXIT_FAILURE); + signal(SIGCHLD, SIG_IGN); + + while (1) { + struct sockaddr_in client_addr; + socklen_t client_len = sizeof(client_addr); + int file_descriptor = + accept(server->file_descriptor, (struct sockaddr *)&client_addr, + &client_len); + if (file_descriptor < 0) { + perror("Could not accept the connection"); + continue; + } + log_message(LOG_LEVEL_INFO, "Received Connection"); + pid_t pid = fork(); + if (pid == 0) { + close(server->file_descriptor); + handler(file_descriptor); + close(file_descriptor); + _exit(0); + } + close(file_descriptor); } - log_message(LOG_LEVEL_INFO, "Received Connection"); - handler(file_descriptor); - close(server->file_descriptor); - close(file_descriptor); } Client *client_create() {