From 8b39a268b9aa3b03ea2db2d98f8f7ea6dcaa12bc Mon Sep 17 00:00:00 2001 From: Vitaliy Filippov Date: Sat, 18 Apr 2026 14:22:33 +0300 Subject: [PATCH] Implement direct AES-256-GCM with a static key for benchmark --- src/client/messenger.cpp | 69 +++++++---- src/client/messenger.h | 12 ++ src/client/msgr_receive.cpp | 180 +++++++++++++++++++++++++++- src/client/msgr_send.cpp | 229 ++++++++++++++++++++++++++++++++++-- src/client/msgr_stop.cpp | 10 ++ 5 files changed, 467 insertions(+), 33 deletions(-) diff --git a/src/client/messenger.cpp b/src/client/messenger.cpp index 194ad2f8..76de6008 100644 --- a/src/client/messenger.cpp +++ b/src/client/messenger.cpp @@ -10,6 +10,7 @@ #include #include "addr_util.h" +#include "str_util.h" #include "messenger.h" #ifdef WITH_RDMA #include "msgr_rdma.h" @@ -308,6 +309,9 @@ void osd_messenger_t::parse_config(const json11::Json & config) osd_tls_ca = config["osd_tls_ca"].string_value(); client_tls_ca = config["client_tls_ca"].string_value(); } + test_osd_aes_key.resize(32); + if (fromhexstr(config["test_osd_aes_key"].string_value(), 32, (uint8_t*)test_osd_aes_key.data()) != 32) + test_osd_aes_key.clear(); if (!osd_num) this->iothread_count = (uint32_t)config["client_iothread_count"].uint64_value(); else @@ -538,10 +542,7 @@ void osd_messenger_t::handle_connect_epoll(int peer_fd) handle_peer_epoll(peer_fd, epoll_events); }); // Check OSD number - if (!tls_cert.empty()) - { - ssl_init(cl, false); - } + ssl_init(cl, false); check_peer_config(cl); } @@ -783,10 +784,7 @@ void osd_messenger_t::accept_connections(int listen_fd) cl->peer_fd = peer_fd; cl->peer_state = PEER_CONNECTED; cl->in_buf = (uint8_t*)malloc_or_die(receive_buffer_size); - if (!tls_cert.empty()) - { - ssl_init(cl, true); - } + ssl_init(cl, true); // Add FD to epoll tfd->set_fd_handler(peer_fd, false, [this](int peer_fd, int epoll_events) { @@ -803,27 +801,48 @@ void osd_messenger_t::accept_connections(int listen_fd) void osd_messenger_t::ssl_init(osd_client_t *cl, bool server_mode) { -#ifdef WITH_OPENSSL - cl->write_to_ssl = BIO_new(BIO_s_mem()); - cl->read_from_ssl = BIO_new(BIO_s_mem()); - cl->ssl_cli = SSL_new(ssl_ctx); - if (!cl->ssl_cli) + if (!tls_cert.empty()) { - fprintf(stderr, "OpenSSL initialization failed: %s\n", ERR_error_string(ERR_get_error(), NULL)); - exit(1); + cl->write_to_ssl = BIO_new(BIO_s_mem()); + cl->read_from_ssl = BIO_new(BIO_s_mem()); + cl->ssl_cli = SSL_new(ssl_ctx); + if (!cl->ssl_cli) + { + fprintf(stderr, "OpenSSL initialization failed: %s\n", ERR_error_string(ERR_get_error(), NULL)); + exit(1); + } + if (server_mode) + { + SSL_set_accept_state(cl->ssl_cli); + } + else + { + SSL_set_connect_state(cl->ssl_cli); + } + SSL_set_bio(cl->ssl_cli, cl->write_to_ssl, cl->read_from_ssl); + bool ok = ssl_do_handshake(cl); + assert(ok); } - if (server_mode) + else if (!test_osd_aes_key.empty()) { - SSL_set_accept_state(cl->ssl_cli); + int r; + cl->enc_ctx = EVP_CIPHER_CTX_new(); + assert(cl->enc_ctx); + r = EVP_EncryptInit_ex(cl->enc_ctx, EVP_aes_256_gcm(), NULL, NULL, NULL); + assert(r == 1); + r = EVP_CIPHER_CTX_set_padding(cl->enc_ctx, 0); + assert(r == 1); + r = EVP_CIPHER_CTX_ctrl(cl->enc_ctx, EVP_CTRL_GCM_SET_IVLEN, 12, NULL); + assert(r == 1); + cl->dec_ctx = EVP_CIPHER_CTX_new(); + assert(cl->dec_ctx); + r = EVP_DecryptInit_ex(cl->dec_ctx, EVP_aes_256_gcm(), NULL, NULL, NULL); + assert(r == 1); + r = EVP_CIPHER_CTX_set_padding(cl->dec_ctx, 0); + assert(r == 1); + r = EVP_CIPHER_CTX_ctrl(cl->dec_ctx, EVP_CTRL_GCM_SET_IVLEN, 12, NULL); + assert(r == 1); } - else - { - SSL_set_connect_state(cl->ssl_cli); - } - SSL_set_bio(cl->ssl_cli, cl->write_to_ssl, cl->read_from_ssl); - bool ok = ssl_do_handshake(cl); - assert(ok); -#endif } #ifdef WITH_RDMA diff --git a/src/client/messenger.h b/src/client/messenger.h index 2d042d7e..60f0d855 100644 --- a/src/client/messenger.h +++ b/src/client/messenger.h @@ -103,6 +103,13 @@ struct osd_client_t msgr_tls_record_hdr_t ssl_read_record; size_t ssl_read_header_size = 0; bool ssl_more_to_buffer = false; + + EVP_CIPHER_CTX *enc_ctx = NULL; + uint8_t enc_tag[16]; + size_t enc_tag_size = 0; + EVP_CIPHER_CTX *dec_ctx = NULL; + uint8_t dec_tag[16]; + size_t dec_tag_size = 0; #endif // Read state @@ -195,9 +202,11 @@ struct __attribute__((visibility("default"))) osd_messenger_t protected: friend class copy_op_reader_t; friend class ssl_op_reader_t; + friend class gcm_op_reader_t; friend class get_op_reader_t; friend class copy_op_writer_t; friend class ssl_op_writer_t; + friend class gcm_op_writer_t; friend class get_op_writer_t; int keepalive_timer_id = -1; @@ -217,6 +226,7 @@ protected: std::string tls_key; std::string osd_tls_ca; std::string client_tls_ca; + std::string test_osd_aes_key; // FIXME Insecure, only for PoC tests #ifdef WITH_RDMA bool use_rdma = true; @@ -328,9 +338,11 @@ protected: bool op_write_buf(osd_client_t *cl, uint8_t *src, size_t src_len, uint8_t *dst, size_t dst_len, bool skip_csum, size_t & from, size_t & done); bool op_copy_data_to(osd_client_t *cl, uint8_t *dst, size_t dst_len, size_t & from, size_t & done); size_t copy_ops_to(osd_client_t *cl, uint8_t *dst, size_t dst_len); + template size_t copy_ops_to_with(osd_client_t *cl, uint8_t *dst, size_t dst_len); void handle_read(int result, osd_client_t *cl); bool handle_read_buffer(osd_client_t *cl, uint8_t *curbuf, size_t bufsize); + template bool handle_buffer_with(osd_client_t *cl, uint8_t *curbuf, size_t bufsize); bool handle_hdr(osd_client_t *cl); bool allocate_op_buffers(osd_client_t *cl); bool allocate_reply_buffers(osd_client_t *cl, osd_op_t *op); diff --git a/src/client/msgr_receive.cpp b/src/client/msgr_receive.cpp index e6577d24..82853a32 100644 --- a/src/client/msgr_receive.cpp +++ b/src/client/msgr_receive.cpp @@ -21,6 +21,7 @@ class msgr_op_reader_t { public: virtual bool read(uint8_t *dst, size_t dst_len, int flags = 0) = 0; + virtual bool finish() = 0; }; class copy_op_reader_t: public msgr_op_reader_t @@ -80,6 +81,11 @@ public: return true; } + bool finish() override + { + return true; + } + size_t get_done() { return done; @@ -252,6 +258,154 @@ buffer_again: return true; } + bool finish() override + { + return true; + } + + size_t get_done() + { + return done; + } +}; + +class gcm_op_reader_t: public msgr_op_reader_t +{ + osd_messenger_t* msgr; + osd_client_t* cl; + size_t from; + + uint8_t *curbuf; + size_t bufsize; + size_t done; + +public: + gcm_op_reader_t(osd_messenger_t* msgr, osd_client_t* cl, uint8_t *curbuf, size_t bufsize): + msgr(msgr), cl(cl), from(cl->read_op_pos), curbuf(curbuf), bufsize(bufsize), done(0) + { + } + + void reset() + { + from = cl->read_op_pos; + uint8_t iv[12] = { 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1 }; + int r = EVP_DecryptInit_ex(cl->dec_ctx, NULL, NULL, (uint8_t*)msgr->test_osd_aes_key.data(), iv); + if (r != 1) + { + fprintf(stderr, "DecryptInit error: "); + ERR_print_errors_fp(stderr); + abort(); + } + } + + bool read(uint8_t *dst, size_t dst_len, int flags) override + { + if (from >= dst_len) + { + // Skip + from -= dst_len; + return true; + } + if (done >= bufsize) + return false; + size_t n = dst_len-from; + if (!(flags & RDR_TLS)) + { + if (n > bufsize-done) + n = bufsize-done; + if (flags & RDR_XTS) + { + msgr->op_decrypted_copy_buf(cl, curbuf, bufsize, dst, dst_len, from, done); + n = 0; + } + else + { + if (cl->read_csum_state && !(flags & RDR_NO_CSUM)) + { + // data may be skipped if dst == NULL but checksum is still calculated + XXH3_64bits_update(cl->read_csum_state, curbuf+done, n); + } + // Here, dst == NULL is allowed + if (dst != NULL) + memcpy(dst+from, curbuf+done, n); + done += n; + } + cl->read_op_pos += n; + from += n; + if (from < dst_len) + { + return false; + } + } + else + { + // Here, dst == NULL is not allowed + assert(dst != NULL); + size_t n = dst_len-from; + if (n > bufsize-done) + n = bufsize-done; + int actual_out; + if (EVP_DecryptUpdate(cl->dec_ctx, dst+from, &actual_out, curbuf+done, n) != 1) + { + fprintf(stderr, "DecryptUpdate error: "); + ERR_print_errors_fp(stderr); + abort(); + } + assert(actual_out == n); + if (cl->read_csum_state && !(flags & RDR_NO_CSUM)) + { + XXH3_64bits_update(cl->read_csum_state, dst+from, n); + } + done += n; + from += n; + cl->read_op_pos += n; + if (from < dst_len) + { + return false; + } + } + from = 0; + return true; + } + + bool finish() override + { + if (cl->dec_tag_size+bufsize-done < 16) + { + // Buffer part of the tag + memcpy(cl->dec_tag+cl->dec_tag_size, curbuf+done, bufsize-done); + cl->dec_tag_size += bufsize-done; + done = bufsize; + return false; + } + int r; + if (cl->dec_tag_size > 0) + { + // Tag is partially buffered, append to it and use it from there + memcpy(cl->dec_tag+cl->dec_tag_size, curbuf+done, 16-cl->dec_tag_size); + done += 16-cl->dec_tag_size; + r = EVP_CIPHER_CTX_ctrl(cl->dec_ctx, EVP_CTRL_GCM_SET_TAG, 16, cl->dec_tag); + } + else + { + // Take full tag directly from the source buffer + r = EVP_CIPHER_CTX_ctrl(cl->dec_ctx, EVP_CTRL_GCM_SET_TAG, 16, curbuf+done); + done += 16; + } + assert(r == 1); + int len = 0; + r = EVP_DecryptFinal_ex(cl->dec_ctx, NULL, &len); + if (r != 1) + { + fprintf(stderr, "Client %ju AES-GCM decryption failed\n", cl->client_id); + cl->io_error = true; + return false; + } + cl->dec_tag_size = 0; + assert(len == 0); + return true; + } + size_t get_done() { return done; @@ -308,6 +462,8 @@ public: bool read(uint8_t *dst, size_t dst_len, int flags) override { + if (cl->dec_ctx) + return false; // FIXME Only for tests, use copy-only with AES if (from >= dst_len) { // Skip @@ -337,6 +493,13 @@ public: from = 0; return true; } + + bool finish() override + { + if (cl->dec_ctx) + return false; + return true; + } }; void osd_messenger_t::read_requests() @@ -529,11 +692,21 @@ void osd_messenger_t::handle_immediate_ops() bool osd_messenger_t::handle_read_buffer(osd_client_t *cl, uint8_t *curbuf, size_t bufsize) { + if (cl->dec_ctx) + { + return handle_buffer_with(cl, curbuf, bufsize); + } + return handle_buffer_with(cl, curbuf, bufsize); +} + +template +bool osd_messenger_t::handle_buffer_with(osd_client_t *cl, uint8_t *curbuf, size_t bufsize) +{ + T rdr(this, cl, curbuf, bufsize); // Reset OSD ping state cl->ping_time_remaining = 0; cl->idle_time_remaining = osd_idle_timeout; // Compose operation(s) from the buffer - ssl_op_reader_t rdr(this, cl, curbuf, bufsize); while (true) { if (!cl->read_op) @@ -802,7 +975,10 @@ bool osd_messenger_t::op_read_from(osd_client_t *cl, msgr_op_reader_t & rdr) if (hdr) { if (!handle_hdr(cl)) + { + cl->io_error = true; return false; + } op = cl->read_op; if (op->op_type == OSD_OP_OUT) goto switched_type; @@ -893,6 +1069,8 @@ switched_type: if (!rdr.read((uint8_t*)&op->csum, 8, RDR_TLS|RDR_NO_CSUM)) return false; } + if (!rdr.finish()) + return false; assert(cl->read_op_pos == cl->read_op_size+OSD_PACKET_SIZE); return true; } diff --git a/src/client/msgr_send.cpp b/src/client/msgr_send.cpp index 40f49bc5..c4b38c2a 100644 --- a/src/client/msgr_send.cpp +++ b/src/client/msgr_send.cpp @@ -23,6 +23,7 @@ class msgr_op_writer_t { public: virtual bool write(uint8_t *src, size_t src_len, int flags = 0) = 0; + virtual bool finish() = 0; }; class copy_op_writer_t: public msgr_op_writer_t @@ -75,6 +76,11 @@ public: return true; } + bool finish() override + { + return true; + } + size_t get_done() { return done; @@ -211,12 +217,150 @@ public: return true; } + bool finish() override + { + return _flush_ssl(); + } + size_t get_done() { return done; } }; +class gcm_op_writer_t: public msgr_op_writer_t +{ + osd_messenger_t* msgr; + osd_client_t* cl; + size_t from; + + uint8_t *curbuf; + size_t bufsize; + size_t done; + +public: + gcm_op_writer_t(osd_messenger_t* msgr, osd_client_t* cl, uint8_t *curbuf, size_t bufsize): + msgr(msgr), cl(cl), from(cl->write_op_pos), curbuf(curbuf), bufsize(bufsize), done(0) + { + } + + void reset() + { + from = cl->write_op_pos; + uint8_t iv[12] = { 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1 }; + int r = EVP_EncryptInit_ex(cl->enc_ctx, NULL, NULL, (uint8_t*)msgr->test_osd_aes_key.data(), iv); + if (r != 1) + { + fprintf(stderr, "EncryptInit error: "); + ERR_print_errors_fp(stderr); + abort(); + } + } + + bool write(uint8_t *src, size_t src_len, int flags) override + { + if (from >= src_len) + { + from -= src_len; + return true; + } + if (!(flags & WR_TLS)) + { + if (flags & WR_XTS) + { + msgr->op_encrypted_copy_buf(cl, curbuf, bufsize, src, src_len, from, done); + } + else + { + size_t n = src_len-from; + if (n > bufsize-done) + n = bufsize-done; + if (!n) + return false; + if (cl->write_csum_state && !(flags & WR_NO_CSUM)) + XXH3_64bits_update(cl->write_csum_state, src+from, n); + memcpy(curbuf+done, src+from, n); + done += n; + cl->write_op_pos += n; + from += n; + } + } + else + { + size_t n = src_len-from; + if (n > bufsize-done) + n = bufsize-done; + if (!n) + return false; + int actual_out; + if (EVP_EncryptUpdate(cl->enc_ctx, curbuf+done, &actual_out, src+from, n) != 1) + { + fprintf(stderr, "EncryptUpdate error: "); + ERR_print_errors_fp(stderr); + abort(); + } + assert(actual_out == n); + if (cl->write_csum_state && !(flags & WR_NO_CSUM)) + XXH3_64bits_update(cl->write_csum_state, src+from, n); + done += n; + cl->write_op_pos += n; + from += n; + } + if (from < src_len) + return false; + from = 0; + return true; + } + + static void write_tag_to(osd_client_t *cl, uint8_t *dst) + { + int actual_out = 0; + int r = EVP_EncryptFinal_ex(cl->enc_ctx, NULL, &actual_out); + if (r != 1) + { + fprintf(stderr, "EncryptFinal error: "); + ERR_print_errors_fp(stderr); + abort(); + } + assert(actual_out == 0); + r = EVP_CIPHER_CTX_ctrl(cl->enc_ctx, EVP_CTRL_GCM_GET_TAG, 16, dst); + assert(r == 1); + } + + bool finish() override + { + // Tag is 16 bytes + if (done >= bufsize) + return false; + if (bufsize-done < 16 || cl->enc_tag_size) + { + // No space for the full tag, but msgr_rdma expects us to always fill the whole buffer + if (!cl->enc_tag_size) + { + write_tag_to(cl, cl->enc_tag); + cl->enc_tag_size = 16; + } + size_t n = bufsize-done; + if (n > cl->enc_tag_size) + n = cl->enc_tag_size; + memcpy(curbuf+done, cl->enc_tag+16-cl->enc_tag_size, n); + done += n; + cl->enc_tag_size -= n; + return !cl->enc_tag_size; + } + // The whole tag fits at once + write_tag_to(cl, curbuf+done); + done += 16; + return true; + } + + size_t get_done() + { + return done; + } +}; + +// FIXME Split into 3 classes - basic, tls and gcm class get_op_writer_t: public msgr_op_writer_t { osd_messenger_t* msgr; @@ -225,9 +369,11 @@ class get_op_writer_t: public msgr_op_writer_t size_t enc_size; size_t done_enc; - void ssl_extend_buf() + void ssl_extend_buf(size_t more = 0) { size_t min_cap = cl->ssl_out_buf_size*2; + if (min_cap < cl->ssl_out_buf_size+more) + min_cap = cl->ssl_out_buf_size+more; if (min_cap < 16384) min_cap = 16384; if (cl->ssl_out_buf_cap < min_cap) @@ -271,6 +417,17 @@ public: from = cl->write_op_pos; enc_size = 0; done_enc = 0; + if (cl->enc_ctx) + { + uint8_t iv[12] = { 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1 }; + int r = EVP_EncryptInit_ex(cl->enc_ctx, NULL, NULL, (uint8_t*)msgr->test_osd_aes_key.data(), iv); + if (r != 1) + { + fprintf(stderr, "EncryptInit error: "); + ERR_print_errors_fp(stderr); + abort(); + } + } } void flush_ssl() @@ -299,9 +456,9 @@ public: { return false; } - if (cl->ssl_cli) + if (flags & WR_TLS) { - if (flags & WR_TLS) + if (cl->ssl_cli) { if (!cl->ssl_handshake_done) { @@ -320,6 +477,30 @@ public: from = 0; return true; } + else if (cl->enc_ctx) + { + // Encrypt data to client's temporary output buffer (all at once) + size_t n = src_len-from; + ssl_extend_buf(n); + int actual_out; + if (EVP_EncryptUpdate(cl->enc_ctx, cl->ssl_out_buf+cl->ssl_out_buf_size, &actual_out, src+from, n) != 1) + { + fprintf(stderr, "EncryptUpdate error: "); + ERR_print_errors_fp(stderr); + abort(); + } + assert(actual_out == n); + if (cl->write_csum_state && !(flags & WR_NO_CSUM)) + XXH3_64bits_update(cl->write_csum_state, src+from, n); + cl->send_list.push_back((iovec){ .iov_base = cl->ssl_out_buf+cl->ssl_out_buf_size, .iov_len = n }); + cl->ssl_out_buf_size += n; + cl->write_op_pos += n; + from += n; + if (from < src_len) + return false; + from = 0; + return true; + } } if (flags & WR_XTS) { @@ -351,6 +532,28 @@ public: from = 0; return true; } + + bool finish() override + { + if (cl->ssl_cli) + { + if (cl->send_list.size() >= IOV_MAX) + return false; + copy_ssl(); + } + else if (cl->enc_ctx) + { + if (cl->send_list.size() >= IOV_MAX) + return false; + // Tag is 16 bytes + ssl_extend_buf(16); + gcm_op_writer_t::write_tag_to(cl, cl->ssl_out_buf+cl->ssl_out_buf_size); + // FIXME coalesce entries in ssl_out_buf + cl->send_list.push_back((iovec){ .iov_base = cl->ssl_out_buf+cl->ssl_out_buf_size, .iov_len = 16 }); + cl->ssl_out_buf_size += 16; + } + return true; + } }; void osd_messenger_t::outbox_push(osd_op_t *cur_op) @@ -539,8 +742,18 @@ bool osd_messenger_t::try_send(osd_client_t *cl) size_t osd_messenger_t::copy_ops_to(osd_client_t *cl, uint8_t *dst, size_t dst_len) { - ssl_op_writer_t wr(this, cl, dst, dst_len); - while ((cl->write_op || cl->write_ops.size()) && wr.get_done() < dst_len) + if (cl->enc_ctx) + { + return copy_ops_to_with(cl, dst, dst_len); + } + return copy_ops_to_with(cl, dst, dst_len); +} + +template +size_t osd_messenger_t::copy_ops_to_with(osd_client_t *cl, uint8_t *dst, size_t dst_len) +{ + T wr(this, cl, dst, dst_len); + while (cl->write_op || cl->write_ops.size()) { if (!cl->write_op) { @@ -560,10 +773,10 @@ size_t osd_messenger_t::copy_ops_to(osd_client_t *cl, uint8_t *dst, size_t dst_l cl->send_free_ops.push_back(op); } } - if (!wr.get_done() && cl->ssl_cli) + /*FIXME if (!wr.get_done() && cl->ssl_cli) { wr.flush_ssl(); - } + }*/ return wr.get_done(); } @@ -783,6 +996,8 @@ bool osd_messenger_t::op_write_to(osd_client_t *cl, msgr_op_writer_t & wr) if (!wr.write((uint8_t*)&cl->write_op->csum, 8, WR_TLS|WR_NO_CSUM)) return false; } + if (!wr.finish()) + return false; op_encrypt_free(cl); cl->write_op = NULL; cl->write_op_pos = 0; diff --git a/src/client/msgr_stop.cpp b/src/client/msgr_stop.cpp index 46ee6e1d..804f2328 100644 --- a/src/client/msgr_stop.cpp +++ b/src/client/msgr_stop.cpp @@ -230,6 +230,16 @@ osd_client_t::~osd_client_t() write_csum_state = NULL; } #ifdef WITH_OPENSSL + if (enc_ctx) + { + EVP_CIPHER_CTX_free(enc_ctx); + enc_ctx = NULL; + } + if (dec_ctx) + { + EVP_CIPHER_CTX_free(dec_ctx); + dec_ctx = NULL; + } if (ssl_cli) { SSL_free(ssl_cli);