diff --git a/src/client/http_client.cpp b/src/client/http_client.cpp index 7768c429..2b838155 100644 --- a/src/client/http_client.cpp +++ b/src/client/http_client.cpp @@ -164,30 +164,42 @@ void http_ares_cb(void *data, ares_socket_t socket_fd, int readable, int writabl } #ifdef WITH_OPENSSL +bool openssl_ctx_add_ca(SSL_CTX *ssl_ctx, const std::string & file_or_pem) +{ + BIO *bio = NULL; + if (file_or_pem.substr(0, 5) != "-----") + { + std::string pem = read_file(file_or_pem); + bio = BIO_new_mem_buf(pem.data(), pem.size()); + } + else + bio = BIO_new_mem_buf(file_or_pem.data(), file_or_pem.size()); + if (!bio) + return false; + X509 *x509 = PEM_read_bio_X509(bio, NULL, 0, NULL); + bool ok = !!x509; + if (x509) + { + X509_STORE *store = SSL_CTX_get_cert_store(ssl_ctx); + X509_STORE_add_cert(store, x509); + X509_free(x509); + } + BIO_free(bio); + return ok; +} + bool openssl_ctx_use_ca(SSL_CTX *ssl_ctx, const std::string & file_or_pem) { if (file_or_pem.substr(0, 5) == "-----") { - BIO *bio = BIO_new_mem_buf(file_or_pem.data(), file_or_pem.size()); - if (!bio) - return false; - X509 *x509 = PEM_read_bio_X509(bio, NULL, 0, NULL); - bool ok = !!x509; - if (x509) - { - X509_STORE *store = SSL_CTX_get_cert_store(ssl_ctx); - X509_STORE_add_cert(store, x509); - X509_free(x509); - } - BIO_free(bio); - return ok; + return openssl_ctx_add_ca(ssl_ctx, file_or_pem); } return file_or_pem.empty() ? !!SSL_CTX_set_default_verify_paths(ssl_ctx) : !!SSL_CTX_load_verify_locations(ssl_ctx, file_or_pem.c_str(), NULL); } -static std::string openssl_get_cn(X509 *x509) +std::string openssl_get_cn(X509 *x509) { X509_NAME* subj = X509_get_subject_name(x509); int pos = X509_NAME_get_index_by_NID(subj, NID_commonName, -1); diff --git a/src/client/http_client.h b/src/client/http_client.h index fc663698..464fb2c5 100644 --- a/src/client/http_client.h +++ b/src/client/http_client.h @@ -8,6 +8,10 @@ #include #include "json11/json11.hpp" +#ifdef WITH_OPENSSL +#include +#endif + #define WS_CONTINUATION 0 #define WS_TEXT 1 #define WS_BINARY 2 @@ -69,3 +73,11 @@ void http_close(http_co_t *co); void http_destroy(http_co_t *co); #pragma GCC visibility pop + +#ifdef WITH_OPENSSL +bool openssl_ctx_add_ca(SSL_CTX *ssl_ctx, const std::string & file_or_pem); +bool openssl_ctx_use_ca(SSL_CTX *ssl_ctx, const std::string & file_or_pem); +std::string openssl_get_cn(X509 *x509); +bool openssl_ctx_use_cert(SSL_CTX *ssl_ctx, const std::string & file_or_pem, std::string & common_name); +bool openssl_ctx_use_key(SSL_CTX *ssl_ctx, const std::string & file_or_pem); +#endif diff --git a/src/client/messenger.cpp b/src/client/messenger.cpp index 89060393..32c666e5 100644 --- a/src/client/messenger.cpp +++ b/src/client/messenger.cpp @@ -14,6 +14,13 @@ #ifdef WITH_RDMA #include "msgr_rdma.h" #endif +#include "http_client.h" +#ifdef WITH_OPENSSL +#include +#include +#include +#include +#endif #include @@ -117,6 +124,43 @@ void msgr_iothread_t::run() void osd_messenger_t::init() { + if (!tls_cert.empty() || !tls_key.empty() || !osd_tls_ca.empty() || !client_tls_ca.empty()) + { + // Initialize TLS context + // FIXME: require OpenSSL +#ifndef WITH_OPENSSL + fprintf(stderr, "Vitastor is built without OpenSSL support\n"); + exit(1); +#else + if (tls_cert.empty() || tls_key.empty() || osd_tls_ca.empty() || osd_num && client_tls_ca.empty()) + { + if (osd_num) + fprintf(stderr, "Vitastor OSD TLS requires osd_tls_cert, osd_tls_key, osd_tls_ca, client_tls_ca\n"); + else + fprintf(stderr, "Vitastor client TLS requires tls_cert, tls_key and osd_tls_ca\n"); + exit(1); + } + else + { + ssl_ctx = SSL_CTX_new(TLS_method()); + SSL_CTX_set_verify(ssl_ctx, SSL_VERIFY_PEER, NULL); + bool ok = SSL_CTX_set_min_proto_version(ssl_ctx, TLS1_3_VERSION); + ok = ok && openssl_ctx_add_ca(ssl_ctx, osd_tls_ca); + if (osd_num) + { + // OSD uses 2 separate root certificates to distinguish between clients and peer OSDs + ok = ok && openssl_ctx_add_ca(ssl_ctx, client_tls_ca); + } + ok = ok && openssl_ctx_use_cert(ssl_ctx, tls_cert, tls_cn); + ok = ok && openssl_ctx_use_key(ssl_ctx, tls_key); + if (!ok) + { + SSL_CTX_free(ssl_ctx); + ssl_ctx = NULL; + } + } +#endif + } #ifdef WITH_RDMACM if (use_rdmacm) { @@ -303,6 +347,13 @@ osd_messenger_t::~osd_messenger_t() { destroy_aes_xts_decrypt(decrypt_ctx); } +#ifdef WITH_OPENSSL + if (ssl_ctx) + { + SSL_CTX_free(ssl_ctx); + ssl_ctx = NULL; + } +#endif } void osd_messenger_t::parse_config(const json11::Json & config) @@ -348,6 +399,19 @@ void osd_messenger_t::parse_config(const json11::Json & config) this->use_proto_checksums = config["proto_checksums"].string_value() == "full" ? MSGR_CSUM_FULL : MSGR_CSUM_PAYLOAD; else this->use_proto_checksums = 0; + if (!osd_num) + { + tls_cert = config["tls_cert"].string_value(); + tls_key = config["tls_key"].string_value(); + osd_tls_ca = config["osd_tls_ca"].string_value(); + } + else + { + tls_cert = config["osd_tls_cert"].string_value(); + tls_key = config["osd_tls_key"].string_value(); + osd_tls_ca = config["osd_tls_ca"].string_value(); + client_tls_ca = config["client_tls_ca"].string_value(); + } if (!osd_num) this->iothread_count = (uint32_t)config["client_iothread_count"].uint64_value(); else @@ -578,6 +642,10 @@ 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); + } check_peer_config(cl); } @@ -819,6 +887,10 @@ 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); + } // Add FD to epoll tfd->set_fd_handler(peer_fd, false, [this](int peer_fd, int epoll_events) { @@ -833,6 +905,26 @@ 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 (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); +#endif +} + #ifdef WITH_RDMA msgr_rdma_context_t* osd_messenger_t::choose_rdma_context(osd_client_t *cl) { diff --git a/src/client/messenger.h b/src/client/messenger.h index 4a095b53..1540886a 100644 --- a/src/client/messenger.h +++ b/src/client/messenger.h @@ -12,6 +12,10 @@ #include #include +#ifdef WITH_OPENSSL +#include +#endif + #include "../util/xxh_x86dispatch.h" #include "../util/robin_hood.h" #include "malloc_or_die.h" @@ -56,6 +60,12 @@ struct op_aes_xts_decrypt_t; void destroy_aes_xts_encrypt(op_aes_xts_encrypt_t *encrypt_ctx); void destroy_aes_xts_decrypt(op_aes_xts_decrypt_t *decrypt_ctx); +struct __attribute__((__packed__)) msgr_tls_record_hdr_t +{ + uint8_t encrypted; + uint32_t size; +}; + struct osd_client_t { uint64_t client_id = 0; @@ -78,6 +88,17 @@ struct osd_client_t msgr_rdma_connection_t *rdma_conn = NULL; #endif +#ifdef WITH_OPENSSL + SSL *ssl_cli = NULL; + BIO *write_to_ssl = NULL; + // FIXME: use custom bio to avoid 1 more memory copy? + BIO *read_from_ssl = NULL; + uint8_t *ssl_out_buf = NULL; + size_t ssl_out_buf_size = 0, ssl_out_buf_cap = 0; + bool ssl_handshake_done = false; + bool ssl_want_write = false; +#endif + // Read state int read_ready = 0; osd_op_t *read_op = NULL; @@ -211,6 +232,11 @@ protected: int iothread_count = 0; int max_aes_xts_pool_size = 256; + std::string tls_cert; + std::string tls_key; + std::string osd_tls_ca; + std::string client_tls_ca; + #ifdef WITH_RDMA bool use_rdma = true; bool use_rdmacm = false; @@ -227,6 +253,19 @@ protected: robin_hood::unordered_flat_map rdmacm_connecting; #endif +#ifdef WITH_OPENSSL + SSL_CTX *ssl_ctx = NULL; + std::string tls_cn; + + void ssl_init(osd_client_t *cl, bool server_mode); + bool ssl_do_handshake(osd_client_t *cl); + bool ssl_do_encrypt(osd_client_t *cl); + size_t ssl_do_encrypt_to(osd_client_t *cl, uint8_t *buf, size_t size); + bool ssl_op_write_buf(osd_client_t *cl, uint8_t *src, size_t src_len, bool skip_csum, size_t & from, size_t & done); + size_t ssl_op_copy_to(osd_client_t *cl, uint8_t *dst, size_t dst_len); + void ssl_op_get_write_buffers(osd_client_t *cl, std::vector & lst); +#endif + std::vector iothreads; std::vector read_ready_clients; std::vector write_ready_clients; @@ -305,6 +344,9 @@ protected: void handle_send(int result, bool prev, bool more, osd_client_t *cl); size_t op_copy_to(osd_client_t *cl, uint8_t *dst, size_t dst_len); void op_get_write_buffers(osd_client_t *cl, std::vector & lst); + void next_write_op(osd_client_t *cl); + 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); void handle_read(int result, osd_client_t *cl); bool handle_read_buffer(osd_client_t *cl, uint8_t *curbuf, size_t bufsize); diff --git a/src/client/msgr_rdma.cpp b/src/client/msgr_rdma.cpp index a9618b24..7ae4e959 100644 --- a/src/client/msgr_rdma.cpp +++ b/src/client/msgr_rdma.cpp @@ -573,11 +573,7 @@ int osd_messenger_t::try_send_rdma_copy(osd_client_t *cl, uint8_t *dst, int dst_ int total_dst_len = dst_len; while (dst_len > 0 && (cl->write_op || cl->write_ops.size())) { - if (!cl->write_op) - { - cl->write_op = cl->write_ops.front(); - cl->write_ops.pop_front(); - } + next_write_op(cl); osd_op_t *op = cl->write_op; size_t copied = op_copy_to(cl, dst, dst_len); if (!copied) diff --git a/src/client/msgr_send.cpp b/src/client/msgr_send.cpp index 6bcea1c4..4b63c2fa 100644 --- a/src/client/msgr_send.cpp +++ b/src/client/msgr_send.cpp @@ -7,6 +7,13 @@ #include "messenger.h" +#ifdef WITH_OPENSSL +#include +#include +#include +#include +#endif + void osd_messenger_t::outbox_push(osd_op_t *cur_op) { assert(cur_op->client_id); @@ -116,25 +123,103 @@ void osd_messenger_t::measure_exec(osd_op_t *cur_op) } } +bool osd_messenger_t::ssl_do_handshake(osd_client_t *cl) +{ + int r = SSL_do_handshake(cl->ssl_cli); + if (r > 0) + { + cl->ssl_handshake_done = true; + } + else + { + r = SSL_get_error(cl->ssl_cli, r); + if (r != SSL_ERROR_WANT_WRITE && r != SSL_ERROR_WANT_READ) + { + fprintf(stderr, "Client %ju TLS handshake error: %s, stopping client\n", cl->client_id, ERR_error_string(ERR_get_error(), NULL)); + stop_client(cl->client_id); + return false; + } + if (r == SSL_ERROR_WANT_WRITE) + { + cl->ssl_want_write = true; + } + } + return true; +} + +bool osd_messenger_t::ssl_do_encrypt(osd_client_t *cl) +{ + if (cl->send_list.size() >= IOV_MAX || !cl->ssl_want_write) + return false; + size_t prev_size = cl->ssl_out_buf_size; + while (true) + { + size_t min_cap = cl->ssl_out_buf_size*2; + if (min_cap < 16384) + min_cap = 16384; + if (cl->ssl_out_buf_cap < min_cap) + { + uint8_t *old_buf = cl->ssl_out_buf; + uint8_t *old_end = old_buf + cl->ssl_out_buf_cap; + cl->ssl_out_buf = (uint8_t*)realloc_or_die(cl->ssl_out_buf, min_cap); + cl->ssl_out_buf_cap = min_cap; + for (auto & iov: cl->send_list) + { + if (iov.iov_base >= old_buf && iov.iov_base < old_end) + iov.iov_base = cl->ssl_out_buf + ((uint8_t*)iov.iov_base - old_buf); + } + } + bool full_read = (ssl_do_encrypt_to(cl, cl->ssl_out_buf+cl->ssl_out_buf_size, + cl->ssl_out_buf_cap-cl->ssl_out_buf_size) == cl->ssl_out_buf_cap-cl->ssl_out_buf_size); + if (!full_read) + break; + } + if (cl->ssl_out_buf_size > prev_size) + cl->send_list.push_back((iovec){ .iov_base = cl->ssl_out_buf+prev_size, .iov_len = cl->ssl_out_buf_size-prev_size }); + return true; +} + +size_t osd_messenger_t::ssl_do_encrypt_to(osd_client_t *cl, uint8_t *buf, size_t size) +{ + if (size < sizeof(msgr_tls_record_hdr_t)) + return 0; + int r = BIO_read(cl->read_from_ssl, buf+sizeof(msgr_tls_record_hdr_t), size-sizeof(msgr_tls_record_hdr_t)); + if (r > 0) + { + if (r < size-sizeof(msgr_tls_record_hdr_t)) + cl->ssl_want_write = false; + msgr_tls_record_hdr_t *hdr = (msgr_tls_record_hdr_t*)buf; + hdr->encrypted = 1; + hdr->size = r; + return r+sizeof(msgr_tls_record_hdr_t); + } + return 0; +} + bool osd_messenger_t::try_send(osd_client_t *cl) { - if (!cl->write_op && !cl->write_ops.size() || cl->write_msg.msg_iovlen > 0 || cl->peer_state == PEER_STOPPED || cl->peer_fd < 0) + if (!cl->write_op && !cl->write_ops.size() && !cl->ssl_want_write || + cl->write_msg.msg_iovlen > 0 || cl->peer_state == PEER_STOPPED || cl->peer_fd < 0) { return true; } assert(cl->peer_state != PEER_RDMA); - while ((cl->write_op || cl->write_ops.size()) && cl->send_list.size() < IOV_MAX) + if (cl->ssl_cli && !cl->ssl_handshake_done) { - if (!cl->write_op) + bool ok = ssl_do_encrypt(cl); + assert(ok && cl->send_list.size() > 0); + } + else + { + while ((cl->write_op || cl->write_ops.size()) && cl->send_list.size() < IOV_MAX) { - cl->write_op = cl->write_ops.front(); - cl->write_ops.pop_front(); - } - osd_op_t *op = cl->write_op; - op_get_write_buffers(cl, cl->send_list); - if (!cl->write_op && op->op_type == OSD_OP_IN) - { - cl->send_free_ops.push_back(op); + next_write_op(cl); + osd_op_t *op = cl->write_op; + op_get_write_buffers(cl, cl->send_list); + if (!cl->write_op && op->op_type == OSD_OP_IN) + { + cl->send_free_ops.push_back(op); + } } } if (ringloop && !use_sync_send_recv) @@ -197,6 +282,21 @@ bool osd_messenger_t::try_send(osd_client_t *cl) return true; } +void osd_messenger_t::next_write_op(osd_client_t *cl) +{ + if (!cl->write_op) + { + cl->write_op = cl->write_ops.front(); + cl->write_ops.pop_front(); + if (cl->proto_csum_status == MSGR_CSUM_FULL || cl->proto_csum_status == MSGR_CSUM_PAYLOAD) + { + if (!cl->write_csum_state) + cl->write_csum_state = XXH3_createState(); + XXH3_64bits_reset(cl->write_csum_state); + } + } +} + void osd_messenger_t::send_replies() { for (int i = 0; i < write_ready_clients.size(); i++) @@ -264,6 +364,7 @@ void osd_messenger_t::handle_send(int result, bool prev, bool more, osd_client_t else delete op; } + cl->ssl_out_buf_size = 0; if (more) cl->zc_free_list.push_back(NULL); // end marker cl->send_free_ops.clear(); @@ -346,20 +447,38 @@ static inline bool op_has_data(osd_op_t *op) op->req.hdr.opcode == OSD_OP_SHOW_CONFIG)) && op->iov.count > 0; } -size_t osd_messenger_t::op_copy_to(osd_client_t *cl, uint8_t *dst, size_t dst_len) +static inline bool op_has_data_for_ssl(osd_op_t *op) { - size_t done = 0; - size_t from = cl->write_op_pos; - auto op_write_buf = [&](uint8_t *src, size_t src_len, bool skip_csum) + return (op->op_type == OSD_OP_IN + ? (op->req.hdr.opcode == OSD_OP_SEC_LIST || + op->req.hdr.opcode == OSD_OP_SHOW_CONFIG || + op->req.hdr.opcode == OSD_OP_DESCRIBE) + : (op->req.hdr.opcode == OSD_OP_SEC_STABILIZE || + op->req.hdr.opcode == OSD_OP_SEC_ROLLBACK || + op->req.hdr.opcode == OSD_OP_SHOW_CONFIG)) && op->iov.count > 0; +} + +static inline bool op_has_data_for_nonssl(osd_op_t *op) +{ + return (op->op_type == OSD_OP_IN + ? (op->req.hdr.opcode == OSD_OP_READ || + op->req.hdr.opcode == OSD_OP_SEC_READ) + : (op->req.hdr.opcode == OSD_OP_WRITE || + op->req.hdr.opcode == OSD_OP_SEC_WRITE || + op->req.hdr.opcode == OSD_OP_SEC_WRITE_STABLE)) && op->iov.count > 0; +} + +bool osd_messenger_t::ssl_op_write_buf(osd_client_t *cl, uint8_t *src, size_t src_len, bool skip_csum, size_t & from, size_t & done) +{ + if (from < src_len) { - if (from < src_len) + size_t n = src_len-from; + int ok = SSL_write_ex(cl->ssl_cli, src+from, n, &n); + if (ok) { - size_t n = src_len-from; - if (n > dst_len-done) - n = dst_len-done; + cl->ssl_want_write = true; if (cl->write_csum_state && !skip_csum) XXH3_64bits_update(cl->write_csum_state, src+from, n); - memcpy(dst+done, src+from, n); done += n; cl->write_op_pos += n; from += n; @@ -368,15 +487,122 @@ size_t osd_messenger_t::op_copy_to(osd_client_t *cl, uint8_t *dst, size_t dst_le from = 0; } else - from -= src_len; - return true; - }; - if ((cl->proto_csum_status == MSGR_CSUM_FULL || cl->proto_csum_status == MSGR_CSUM_PAYLOAD) && !from) - { - if (!cl->write_csum_state) - cl->write_csum_state = XXH3_createState(); - XXH3_64bits_reset(cl->write_csum_state); + { + int res = SSL_get_error(cl->ssl_cli, ok); + if (res == SSL_ERROR_WANT_WRITE || res == 0) + cl->ssl_want_write = true; + else if (res == SSL_ERROR_ZERO_RETURN) + { + fprintf(stderr, "Client %ju TLS disconnected\n", cl->client_id); + stop_client(cl->client_id); + } + else if (res != SSL_ERROR_WANT_READ) + { + fprintf(stderr, "Client %ju TLS write error: %s. Disconnecting client\n", cl->client_id, ERR_error_string(ERR_get_error(), NULL)); + stop_client(cl->client_id); + } + return false; + } } + else + from -= src_len; + return true; +} + +bool osd_messenger_t::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) +{ + if (from < src_len) + { + size_t n = src_len-from; + if (n > dst_len-done) + n = dst_len-done; + if (cl->write_csum_state && !skip_csum) + XXH3_64bits_update(cl->write_csum_state, src+from, n); + memcpy(dst+done, src+from, n); + done += n; + cl->write_op_pos += n; + from += n; + if (from < src_len) + return false; + from = 0; + } + else + from -= src_len; + return true; +} + +size_t osd_messenger_t::ssl_op_copy_to(osd_client_t *cl, uint8_t *dst, size_t dst_len) +{ + size_t done = 0; + size_t from = cl->write_op_pos; + // Encrypt headers and data except read/write data + auto to_ssl = [&](uint8_t *src, size_t src_len, bool skip_csum) + { + return ssl_op_write_buf(cl, src, src_len, skip_csum, from, done); + }; + int i = 0; + bool full_hdr = false; + bool full_op = false; + do + { + if (!full_hdr) + full_hdr = op_write_headers(cl->write_op, to_ssl, cl->proto_csum_status != MSGR_CSUM_FULL); + if (full_hdr && op_has_data_for_ssl(cl->write_op)) + { + for (; i < cl->write_op->iov.count; i++) + if (!ssl_op_write_buf(cl, (uint8_t*)cl->write_op->iov.buf[i].iov_base, cl->write_op->iov.buf[i].iov_len, false, from, done)) + break; + full_op = true; + } + if (cl->peer_state == PEER_STOPPED) + return 0; + if (cl->ssl_want_write) + { + auto ssl_written = ssl_do_encrypt_to(cl, dst+done, dst_len-done); + if (!ssl_written) + return done; + done += ssl_written; + } + } while (!full_op); + // Non-TLS-encrypted operation data + if (op_has_data_for_nonssl(cl->write_op)) + { + if (!op_copy_data_to(cl, dst, dst_len, from, done)) + return done; + } + // TLS-encrypted checksum (uh oh...) + if (cl->proto_csum_status == MSGR_CSUM_FULL || + cl->proto_csum_status == MSGR_CSUM_PAYLOAD && cl->write_op_pos > OSD_PACKET_SIZE) + { + if (!from) + cl->write_op->csum = XXH3_64bits_digest(cl->write_csum_state); + if (!ssl_op_write_buf(cl, (uint8_t*)&cl->write_op->csum, 8, true, from, done)) + return done; + if (cl->ssl_want_write) + { + auto ssl_written = ssl_do_encrypt_to(cl, dst+done, dst_len-done); + if (!ssl_written) + return done; + done += ssl_written; + } + } + cl->write_op = NULL; + cl->write_op_pos = 0; + return done; +} + +size_t osd_messenger_t::op_copy_to(osd_client_t *cl, uint8_t *dst, size_t dst_len) +{ + if (cl->ssl_cli) + { + return ssl_op_copy_to(cl, dst, dst_len); + } + size_t done = 0; + size_t from = cl->write_op_pos; + auto op_write_buf = [&](uint8_t *src, size_t src_len, bool skip_csum) + { + return this->op_write_buf(cl, src, src_len, dst, dst_len, skip_csum, from, done); + }; // Header if (!op_write_headers(cl->write_op, op_write_buf, cl->proto_csum_status != MSGR_CSUM_FULL)) { @@ -385,21 +611,8 @@ size_t osd_messenger_t::op_copy_to(osd_client_t *cl, uint8_t *dst, size_t dst_le // Operation data if (op_has_data(cl->write_op)) { - if (cl->write_op->enc) - { - if (!op_encrypted_copy_data_to(cl, dst, dst_len, from, done)) - { - return done; - } - } - else - { - for (int i = 0; i < cl->write_op->iov.count; i++) - { - if (!op_write_buf((uint8_t*)cl->write_op->iov.buf[i].iov_base, cl->write_op->iov.buf[i].iov_len, false)) - return done; - } - } + if (!op_copy_data_to(cl, dst, dst_len, from, done)) + return done; } if (cl->proto_csum_status == MSGR_CSUM_FULL || cl->proto_csum_status == MSGR_CSUM_PAYLOAD && cl->write_op_pos > OSD_PACKET_SIZE) @@ -414,6 +627,27 @@ size_t osd_messenger_t::op_copy_to(osd_client_t *cl, uint8_t *dst, size_t dst_le return done; } +bool osd_messenger_t::op_copy_data_to(osd_client_t *cl, uint8_t *dst, size_t dst_len, size_t & from, size_t & done) +{ + if (cl->write_op->enc) + { + if (!op_encrypted_copy_data_to(cl, dst, dst_len, from, done)) + { + return false; + } + } + else + { + for (int i = 0; i < cl->write_op->iov.count; i++) + { + auto & iov = cl->write_op->iov.buf[i]; + if (!op_write_buf(cl, (uint8_t*)iov.iov_base, iov.iov_len, dst, dst_len, false, from, done)) + return false; + } + } + return true; +} + void osd_messenger_t::op_get_write_buffers(osd_client_t *cl, std::vector & lst) { size_t from = cl->write_op_pos; @@ -433,12 +667,6 @@ void osd_messenger_t::op_get_write_buffers(osd_client_t *cl, std::vector from -= src_len; return true; }; - if ((cl->proto_csum_status == MSGR_CSUM_FULL || cl->proto_csum_status == MSGR_CSUM_PAYLOAD) && !from) - { - if (!cl->write_csum_state) - cl->write_csum_state = XXH3_createState(); - XXH3_64bits_reset(cl->write_csum_state); - } // Header if (!op_write_headers(cl->write_op, op_write_buf, cl->proto_csum_status != MSGR_CSUM_FULL)) { @@ -482,3 +710,89 @@ void osd_messenger_t::op_get_write_buffers(osd_client_t *cl, std::vector cl->write_op = NULL; cl->write_op_pos = 0; } + +void osd_messenger_t::ssl_op_get_write_buffers(osd_client_t *cl, std::vector & lst) +{ + size_t done = 0; + size_t from = cl->write_op_pos; + auto op_write_buf = [&](uint8_t *src, size_t src_len, bool skip_csum) + { + if (lst.size() >= IOV_MAX) + return false; + if (from < src_len) + { + if (cl->write_csum_state && !skip_csum) + XXH3_64bits_update(cl->write_csum_state, src+from, src_len-from); + lst.push_back((iovec){ .iov_base = src+from, .iov_len = src_len-from }); + cl->write_op_pos += src_len-from; + from = 0; + } + else + from -= src_len; + return true; + }; + // Encrypt headers and data except read/write data + auto to_ssl = [&](uint8_t *src, size_t src_len, bool skip_csum) + { + return ssl_op_write_buf(cl, src, src_len, skip_csum, from, done); + }; + int i = 0; + bool full_hdr = false; + bool full_op = false; + do + { + if (!full_hdr) + full_hdr = op_write_headers(cl->write_op, to_ssl, cl->proto_csum_status != MSGR_CSUM_FULL); + if (full_hdr && op_has_data_for_ssl(cl->write_op)) + { + for (; i < cl->write_op->iov.count; i++) + if (!ssl_op_write_buf(cl, (uint8_t*)cl->write_op->iov.buf[i].iov_base, cl->write_op->iov.buf[i].iov_len, false, from, done)) + break; + full_op = true; + } + if (cl->peer_state == PEER_STOPPED) + return; + if (!ssl_do_encrypt(cl)) + return; + } while (!full_op); + // Non-TLS-encrypted operation data + if (op_has_data_for_nonssl(cl->write_op)) + { + if (cl->write_op->enc) + { + if (lst.size() >= IOV_MAX) + return; + // No way except to allocate a temporary buffer and encrypt data to it + assert(cl->write_op->req.hdr.opcode == OSD_OP_WRITE); + size_t remsize = cl->write_op->req.rw.len - from + (from % 16); + assert(remsize > 0); + assert(!cl->write_op->enc_buf); + cl->write_op->enc_buf = (uint8_t*)malloc_or_die(remsize); + size_t done = 0; + bool end = op_encrypted_copy_data_to(cl, cl->write_op->enc_buf, remsize, from, done); + assert(end); + lst.push_back((iovec){ .iov_base = cl->write_op->enc_buf, .iov_len = remsize }); + } + else + { + for (int i = 0; i < cl->write_op->iov.count; i++) + { + if (!op_write_buf((uint8_t*)cl->write_op->iov.buf[i].iov_base, cl->write_op->iov.buf[i].iov_len, false)) + return; + } + } + } + // TLS-encrypted checksum (uh oh...) + if (cl->proto_csum_status == MSGR_CSUM_FULL || + cl->proto_csum_status == MSGR_CSUM_PAYLOAD && cl->write_op_pos > OSD_PACKET_SIZE) + { + if (!from) + cl->write_op->csum = XXH3_64bits_digest(cl->write_csum_state); + if (!ssl_op_write_buf(cl, (uint8_t*)&cl->write_op->csum, 8, true, from, done)) + return; + if (!ssl_do_encrypt(cl)) + return; + } + cl->write_op = NULL; + cl->write_op_pos = 0; +} diff --git a/src/client/msgr_stop.cpp b/src/client/msgr_stop.cpp index bd7f6d07..c73e54ba 100644 --- a/src/client/msgr_stop.cpp +++ b/src/client/msgr_stop.cpp @@ -9,6 +9,12 @@ #ifdef WITH_RDMA #include "msgr_rdma.h" #endif +#ifdef WITH_OPENSSL +#include +#include +#include +#include +#endif void osd_client_t::cancel_ops() { @@ -239,4 +245,18 @@ osd_client_t::~osd_client_t() XXH3_freeState(write_csum_state); write_csum_state = NULL; } +#ifdef WITH_OPENSSL + if (ssl_cli) + { + SSL_free(ssl_cli); + ssl_cli = NULL; + write_to_ssl = NULL; + read_from_ssl = NULL; + } + if (ssl_out_buf) + { + free(ssl_out_buf); + ssl_out_buf = NULL; + } +#endif }