From 8b8cfbe1195bef6384674641225d40b5e98e88f4 Mon Sep 17 00:00:00 2001 From: matt335672 <30179339+matt335672@users.noreply.github.com> Date: Thu, 27 Jan 2022 16:31:53 +0000 Subject: [PATCH] Improve wrapping of openssl module --- common/ssl_calls.c | 26 ++++++++++++++++++++++---- common/ssl_calls.h | 22 ++++++++-------------- common/trans.c | 12 ++++-------- 3 files changed, 34 insertions(+), 26 deletions(-) diff --git a/common/ssl_calls.c b/common/ssl_calls.c index 792a9130..168add76 100644 --- a/common/ssl_calls.c +++ b/common/ssl_calls.c @@ -53,6 +53,17 @@ static EVP_CIPHER *g_cipher_des_ede3_cbc; /* DES3 CBC cipher */ static EVP_MAC *g_mac_hmac; /* HMAC MAC */ #endif +/* definition of ssl_tls */ +struct ssl_tls +{ + SSL *ssl; /* SSL * */ + SSL_CTX *ctx; /* SSL_CTX * */ + char *cert; + char *key; + struct trans *trans; + tintptr rwo; /* wait obj */ + int error_logged; /* Error has already been logged */ +}; #if OPENSSL_VERSION_NUMBER < 0x10100000L static inline HMAC_CTX * @@ -1392,16 +1403,23 @@ ssl_tls_can_recv(struct ssl_tls *tls, int sck, int millis) /*****************************************************************************/ const char * -ssl_get_version(const struct ssl_st *ssl) +ssl_get_version(const struct ssl_tls *ssl) { - return SSL_get_version(ssl); + return SSL_get_version(ssl->ssl); } /*****************************************************************************/ const char * -ssl_get_cipher_name(const struct ssl_st *ssl) +ssl_get_cipher_name(const struct ssl_tls *ssl) { - return SSL_get_cipher_name(ssl); + return SSL_get_cipher_name(ssl->ssl); +} + +/*****************************************************************************/ +tintptr +ssl_get_rwo(const struct ssl_tls *ssl) +{ + return ssl->rwo; } /*****************************************************************************/ diff --git a/common/ssl_calls.h b/common/ssl_calls.h index 78edd946..01b67406 100644 --- a/common/ssl_calls.h +++ b/common/ssl_calls.h @@ -22,6 +22,10 @@ #include "arch.h" +/* Incomplete types */ +struct ssl_tls; +struct trans; + int ssl_init(void); int @@ -81,18 +85,6 @@ int ssl_gen_key_xrdp1(int key_size_in_bits, const char *exp, int exp_len, char *mod, int mod_len, char *pri, int pri_len); -/* ssl_tls */ -struct ssl_tls -{ - struct ssl_st *ssl; /* SSL * */ - struct ssl_ctx_st *ctx; /* SSL_CTX * */ - char *cert; - char *key; - struct trans *trans; - tintptr rwo; /* wait obj */ - int error_logged; /* Error has already been logged */ -}; - /* xrdp_tls.c */ struct ssl_tls * ssl_tls_create(struct trans *trans, const char *key, const char *cert); @@ -110,12 +102,14 @@ ssl_tls_write(struct ssl_tls *tls, const char *data, int length); int ssl_tls_can_recv(struct ssl_tls *tls, int sck, int millis); const char * -ssl_get_version(const struct ssl_st *ssl); +ssl_get_version(const struct ssl_tls *ssl); const char * -ssl_get_cipher_name(const struct ssl_st *ssl); +ssl_get_cipher_name(const struct ssl_tls *ssl); int ssl_get_protocols_from_string(const char *str, long *ssl_protocols); const char * get_openssl_version(); +tintptr +ssl_get_rwo(const struct ssl_tls *ssl); #endif diff --git a/common/trans.c b/common/trans.c index 3afae55d..55d2a638 100644 --- a/common/trans.c +++ b/common/trans.c @@ -179,13 +179,9 @@ trans_get_wait_objs(struct trans *self, tbus *objs, int *count) objs[*count] = self->sck; (*count)++; - if (self->tls != 0) + if (self->tls != NULL && (objs[*count] = ssl_get_rwo(self->tls)) != 0) { - if (self->tls->rwo != 0) - { - objs[*count] = self->tls->rwo; - (*count)++; - } + (*count)++; } return 0; @@ -995,8 +991,8 @@ trans_set_tls_mode(struct trans *self, const char *key, const char *cert, self->trans_send = trans_tls_send; self->trans_can_recv = trans_tls_can_recv; - self->ssl_protocol = ssl_get_version(self->tls->ssl); - self->cipher_name = ssl_get_cipher_name(self->tls->ssl); + self->ssl_protocol = ssl_get_version(self->tls); + self->cipher_name = ssl_get_cipher_name(self->tls); return 0; }