Merge branch 'issue326-add-openssl-error-handler-callback' of https://github.com/tbeu/paho.mqtt.c into develop

This commit is contained in:
Ian Craggs 2018-08-31 14:52:28 +01:00
commit e2fe35c53c
8 changed files with 129 additions and 38 deletions

View File

@ -2702,7 +2702,7 @@ int MQTTAsync_connect(MQTTAsync handle, const MQTTAsync_connectOptions* options)
}
if (options->struct_version != 0 && options->ssl) /* check validity of SSL options structure */
{
if (strncmp(options->ssl->struct_id, "MQTS", 4) != 0 || options->ssl->struct_version < 0 || options->ssl->struct_version > 2)
if (strncmp(options->ssl->struct_id, "MQTS", 4) != 0 || options->ssl->struct_version < 0 || options->ssl->struct_version > 3)
{
rc = MQTTASYNC_BAD_STRUCTURE;
goto exit;
@ -2859,6 +2859,11 @@ int MQTTAsync_connect(MQTTAsync handle, const MQTTAsync_connectOptions* options)
if (options->ssl->CApath)
m->c->sslopts->CApath = MQTTStrdup(options->ssl->CApath);
}
if (m->c->sslopts->struct_version >= 3)
{
m->c->sslopts->ssl_error_cb = options->ssl->ssl_error_cb;
m->c->sslopts->ssl_error_context = options->ssl->ssl_error_context;
}
}
#else
if (options->struct_version != 0 && options->ssl)
@ -3448,8 +3453,11 @@ static int MQTTAsync_connecting(MQTTAsyncs* m)
if (m->c->session != NULL)
if ((rc = SSL_set_session(m->c->net.ssl, m->c->session)) != 1)
Log(TRACE_MIN, -1, "Failed to set SSL session with stored data, non critical");
rc = SSLSocket_connect(m->c->net.ssl, m->c->net.socket,
m->serverURI, m->c->sslopts->verify);
rc = m->c->sslopts->struct_version >= 3 ?
SSLSocket_connect(m->c->net.ssl, m->c->net.socket, m->serverURI,
m->c->sslopts->verify, m->c->sslopts->ssl_error_cb, m->c->sslopts->ssl_error_context) :
SSLSocket_connect(m->c->net.ssl, m->c->net.socket, m->serverURI,
m->c->sslopts->verify, NULL, NULL);
if (rc == TCPSOCKET_INTERRUPTED)
{
rc = MQTTCLIENT_SUCCESS; /* the connect is still in progress */
@ -3512,8 +3520,12 @@ static int MQTTAsync_connecting(MQTTAsyncs* m)
#if defined(OPENSSL)
else if (m->c->connect_state == SSL_IN_PROGRESS) /* SSL connect sent - wait for completion */
{
if ((rc = SSLSocket_connect(m->c->net.ssl, m->c->net.socket,
m->serverURI, m->c->sslopts->verify)) != 1)
rc = m->c->sslopts->struct_version >= 3 ?
SSLSocket_connect(m->c->net.ssl, m->c->net.socket, m->serverURI,
m->c->sslopts->verify, m->c->sslopts->ssl_error_cb, m->c->sslopts->ssl_error_context) :
SSLSocket_connect(m->c->net.ssl, m->c->net.socket, m->serverURI,
m->c->sslopts->verify, NULL, NULL);
if (rc != 1)
goto exit;
if(!m->c->cleansession && m->c->session == NULL)

View File

@ -944,9 +944,21 @@ typedef struct
* Exists only if struct_version >= 2
*/
const char* CApath;
/**
* Callback function for OpenSSL error handler ERR_print_errors_cb
* Exists only if struct_version >= 3
*/
int (*ssl_error_cb) (const char *str, size_t len, void *u);
/**
* Application-specific contex for OpenSSL error handler ERR_print_errors_cb
* Exists only if struct_version >= 3
*/
void* ssl_error_context;
} MQTTAsync_SSLOptions;
#define MQTTAsync_SSLOptions_initializer { {'M', 'Q', 'T', 'S'}, 2, NULL, NULL, NULL, NULL, NULL, 1, MQTT_SSL_VERSION_DEFAULT, 0, NULL }
#define MQTTAsync_SSLOptions_initializer { {'M', 'Q', 'T', 'S'}, 3, NULL, NULL, NULL, NULL, NULL, 1, MQTT_SSL_VERSION_DEFAULT, 0, NULL, NULL, NULL }
/**
* MQTTAsync_connectOptions defines several settings that control the way the

View File

@ -882,8 +882,11 @@ static thread_return_type WINAPI MQTTClient_run(void* n)
#if defined(OPENSSL)
else if (m->c->connect_state == SSL_IN_PROGRESS)
{
rc = SSLSocket_connect(m->c->net.ssl, m->c->net.socket,
m->serverURI, m->c->sslopts->verify);
rc = m->c->sslopts->struct_version >= 3 ?
SSLSocket_connect(m->c->net.ssl, m->c->net.socket, m->serverURI,
m->c->sslopts->verify, m->c->sslopts->ssl_error_cb, m->c->sslopts->ssl_error_context) :
SSLSocket_connect(m->c->net.ssl, m->c->net.socket, m->serverURI,
m->c->sslopts->verify, NULL, NULL);
if (rc == 1 || rc == SSL_FATAL)
{
if (rc == 1 && (m->c->cleansession == 0 && m->c->cleanstart == 0) && m->c->session == NULL)
@ -1138,8 +1141,11 @@ static MQTTResponse MQTTClient_connectURIVersion(MQTTClient handle, MQTTClient_c
if (m->c->session != NULL)
if ((rc = SSL_set_session(m->c->net.ssl, m->c->session)) != 1)
Log(TRACE_MIN, -1, "Failed to set SSL session with stored data, non critical");
rc = SSLSocket_connect(m->c->net.ssl, m->c->net.socket,
m->serverURI, m->c->sslopts->verify);
rc = m->c->sslopts->struct_version >= 3 ?
SSLSocket_connect(m->c->net.ssl, m->c->net.socket, m->serverURI,
m->c->sslopts->verify, m->c->sslopts->ssl_error_cb, m->c->sslopts->ssl_error_context) :
SSLSocket_connect(m->c->net.ssl, m->c->net.socket, m->serverURI,
m->c->sslopts->verify, NULL, NULL);
if (rc == TCPSOCKET_INTERRUPTED)
m->c->connect_state = SSL_IN_PROGRESS; /* the connect is still in progress */
else if (rc == SSL_FATAL)
@ -1438,6 +1444,11 @@ static MQTTResponse MQTTClient_connectURI(MQTTClient handle, MQTTClient_connectO
if (options->ssl->CApath)
m->c->sslopts->CApath = MQTTStrdup(options->ssl->CApath);
}
if (m->c->sslopts->struct_version >= 3)
{
m->c->sslopts->ssl_error_cb = options->ssl->ssl_error_cb;
m->c->sslopts->ssl_error_context = options->ssl->ssl_error_context;
}
}
#endif
@ -1535,7 +1546,7 @@ MQTTResponse MQTTClient_connect5(MQTTClient handle, MQTTClient_connectOptions* o
#if defined(OPENSSL)
if (options->struct_version != 0 && options->ssl) /* check validity of SSL options structure */
{
if (strncmp(options->ssl->struct_id, "MQTS", 4) != 0 || options->ssl->struct_version < 0 || options->ssl->struct_version > 2)
if (strncmp(options->ssl->struct_id, "MQTS", 4) != 0 || options->ssl->struct_version < 0 || options->ssl->struct_version > 3)
{
rc.reasonCode = MQTTCLIENT_BAD_STRUCTURE;
goto exit;
@ -2381,8 +2392,11 @@ static MQTTPacket* MQTTClient_waitfor(MQTTClient handle, int packet_type, int* r
#if defined(OPENSSL)
else if (m->c->connect_state == SSL_IN_PROGRESS)
{
*rc = SSLSocket_connect(m->c->net.ssl, sock,
m->serverURI, m->c->sslopts->verify);
*rc = m->c->sslopts->struct_version >= 3 ?
SSLSocket_connect(m->c->net.ssl, sock, m->serverURI,
m->c->sslopts->verify, m->c->sslopts->ssl_error_cb, m->c->sslopts->ssl_error_context) :
SSLSocket_connect(m->c->net.ssl, sock, m->serverURI,
m->c->sslopts->verify, NULL, NULL);
if (*rc == SSL_FATAL)
break;
else if (*rc == 1) /* rc == 1 means SSL connect has finished and succeeded */

View File

@ -694,9 +694,21 @@ typedef struct
*/
const char* CApath;
/**
* Callback function for OpenSSL error handler ERR_print_errors_cb
* Exists only if struct_version >= 3
*/
int (*ssl_error_cb) (const char *str, size_t len, void *u);
/**
* Application-specific contex for OpenSSL error handler ERR_print_errors_cb
* Exists only if struct_version >= 3
*/
void* ssl_error_context;
} MQTTClient_SSLOptions;
#define MQTTClient_SSLOptions_initializer { {'M', 'Q', 'T', 'S'}, 2, NULL, NULL, NULL, NULL, NULL, 1, MQTT_SSL_VERSION_DEFAULT, 0, NULL }
#define MQTTClient_SSLOptions_initializer { {'M', 'Q', 'T', 'S'}, 3, NULL, NULL, NULL, NULL, NULL, 1, MQTT_SSL_VERSION_DEFAULT, 0, NULL, NULL, NULL }
/**
* MQTTClient_connectOptions defines several settings that control the way the

View File

@ -125,8 +125,11 @@ int MQTTProtocol_connect(const char* ip_address, Clients* aClient, int websocket
{
if (SSLSocket_setSocketForSSL(&aClient->net, aClient->sslopts, ip_address, addr_len) == 1)
{
rc = SSLSocket_connect(aClient->net.ssl, aClient->net.socket,
ip_address, aClient->sslopts->verify);
rc = aClient->sslopts->struct_version >= 3 ?
SSLSocket_connect(aClient->net.ssl, aClient->net.socket, ip_address,
aClient->sslopts->verify, aClient->sslopts->ssl_error_cb, aClient->sslopts->ssl_error_context) :
SSLSocket_connect(aClient->net.ssl, aClient->net.socket, ip_address,
aClient->sslopts->verify, NULL, NULL);
if (rc == TCPSOCKET_INTERRUPTED)
aClient->connect_state = SSL_IN_PROGRESS; /* SSL connect called - wait for completion */
}

View File

@ -46,7 +46,7 @@
extern Sockets s;
int SSLSocket_error(char* aString, SSL* ssl, int sock, int rc);
static int SSLSocket_error(char* aString, SSL* ssl, int sock, int rc, int (*cb)(const char *str, size_t len, void *u), void* u);
char* SSL_get_verify_result_string(int rc);
void SSL_CTX_info_callback(const SSL* ssl, int where, int ret);
char* SSLSocket_get_version_string(int version);
@ -85,9 +85,12 @@ static ssl_mutex_type sslCoreMutex;
* Gets the specific error corresponding to SOCKET_ERROR
* @param aString the function that was being used when the error occurred
* @param sock the socket on which the error occurred
* @param rc the return code
* @param cb the callback function to be passed as first argument to ERR_print_errors_cb
* @param u context to be passed as second argument to ERR_print_errors_cb
* @return the specific TCP error code
*/
int SSLSocket_error(char* aString, SSL* ssl, int sock, int rc)
static int SSLSocket_error(char* aString, SSL* ssl, int sock, int rc, int (*cb)(const char *str, size_t len, void *u), void* u)
{
int error;
@ -106,7 +109,8 @@ int SSLSocket_error(char* aString, SSL* ssl, int sock, int rc)
if (strcmp(aString, "shutdown") != 0)
Log(TRACE_MIN, -1, "SSLSocket error %s(%d) in %s for socket %d rc %d errno %d %s\n", buf, error, aString, sock, rc, errno, strerror(errno));
ERR_print_errors_fp(stderr);
if (cb)
ERR_print_errors_cb(cb, u);
if (error == SSL_ERROR_SSL || error == SSL_ERROR_SYSCALL)
error = SSL_FATAL;
}
@ -548,7 +552,10 @@ int SSLSocket_createContext(networkHandles* net, MQTTClient_SSLOptions* opts)
}
if (net->ctx == NULL)
{
SSLSocket_error("SSL_CTX_new", NULL, net->socket, rc);
if (opts->struct_version >= 3)
SSLSocket_error("SSL_CTX_new", NULL, net->socket, rc, opts->ssl_error_cb, opts->ssl_error_context);
else
SSLSocket_error("SSL_CTX_new", NULL, net->socket, rc, NULL, NULL);
goto exit;
}
}
@ -557,7 +564,10 @@ int SSLSocket_createContext(networkHandles* net, MQTTClient_SSLOptions* opts)
{
if ((rc = SSL_CTX_use_certificate_chain_file(net->ctx, opts->keyStore)) != 1)
{
SSLSocket_error("SSL_CTX_use_certificate_chain_file", NULL, net->socket, rc);
if (opts->struct_version >= 3)
SSLSocket_error("SSL_CTX_use_certificate_chain_file", NULL, net->socket, rc, opts->ssl_error_cb, opts->ssl_error_context);
else
SSLSocket_error("SSL_CTX_use_certificate_chain_file", NULL, net->socket, rc, NULL, NULL);
goto free_ctx; /*If we can't load the certificate (chain) file then loading the privatekey won't work either as it needs a matching cert already loaded */
}
@ -576,7 +586,10 @@ int SSLSocket_createContext(networkHandles* net, MQTTClient_SSLOptions* opts)
opts->privateKey = NULL;
if (rc != 1)
{
SSLSocket_error("SSL_CTX_use_PrivateKey_file", NULL, net->socket, rc);
if (opts->struct_version >= 3)
SSLSocket_error("SSL_CTX_use_PrivateKey_file", NULL, net->socket, rc, opts->ssl_error_cb, opts->ssl_error_context);
else
SSLSocket_error("SSL_CTX_use_PrivateKey_file", NULL, net->socket, rc, NULL, NULL);
goto free_ctx;
}
}
@ -585,13 +598,19 @@ int SSLSocket_createContext(networkHandles* net, MQTTClient_SSLOptions* opts)
{
if ((rc = SSL_CTX_load_verify_locations(net->ctx, opts->trustStore, NULL)) != 1)
{
SSLSocket_error("SSL_CTX_load_verify_locations", NULL, net->socket, rc);
if (opts->struct_version >= 3)
SSLSocket_error("SSL_CTX_load_verify_locations", NULL, net->socket, rc, opts->ssl_error_cb, opts->ssl_error_context);
else
SSLSocket_error("SSL_CTX_load_verify_locations", NULL, net->socket, rc, NULL, NULL);
goto free_ctx;
}
}
else if ((rc = SSL_CTX_set_default_verify_paths(net->ctx)) != 1)
{
SSLSocket_error("SSL_CTX_set_default_verify_paths", NULL, net->socket, rc);
if (opts->struct_version >= 3)
SSLSocket_error("SSL_CTX_set_default_verify_paths", NULL, net->socket, rc, opts->ssl_error_cb, opts->ssl_error_context);
else
SSLSocket_error("SSL_CTX_set_default_verify_paths", NULL, net->socket, rc, NULL, NULL);
goto free_ctx;
}
@ -599,7 +618,10 @@ int SSLSocket_createContext(networkHandles* net, MQTTClient_SSLOptions* opts)
{
if ((rc = SSL_CTX_set_cipher_list(net->ctx, opts->enabledCipherSuites)) != 1)
{
SSLSocket_error("SSL_CTX_set_cipher_list", NULL, net->socket, rc);
if (opts->struct_version >= 3)
SSLSocket_error("SSL_CTX_set_cipher_list", NULL, net->socket, rc, opts->ssl_error_cb, opts->ssl_error_context);
else
SSLSocket_error("SSL_CTX_set_cipher_list", NULL, net->socket, rc, NULL, NULL);
goto free_ctx;
}
}
@ -643,14 +665,21 @@ int SSLSocket_setSocketForSSL(networkHandles* net, MQTTClient_SSLOptions* opts,
if (cipher == NULL)
break;
Log(TRACE_PROTOCOL, 1, "SSL cipher available: %d:%s", i, cipher);
}
if ((rc = SSL_set_fd(net->ssl, net->socket)) != 1)
SSLSocket_error("SSL_set_fd", net->ssl, net->socket, rc);
}
if ((rc = SSL_set_fd(net->ssl, net->socket)) != 1) {
if (opts->struct_version >= 3)
SSLSocket_error("SSL_set_fd", net->ssl, net->socket, rc, opts->ssl_error_cb, opts->ssl_error_context);
else
SSLSocket_error("SSL_set_fd", net->ssl, net->socket, rc, NULL, NULL);
}
hostname_plus_null = malloc(hostname_len + 1u );
MQTTStrncpy(hostname_plus_null, hostname, hostname_len + 1u);
if ((rc = SSL_set_tlsext_host_name(net->ssl, hostname_plus_null)) != 1)
SSLSocket_error("SSL_set_tlsext_host_name", NULL, net->socket, rc);
if ((rc = SSL_set_tlsext_host_name(net->ssl, hostname_plus_null)) != 1) {
if (opts->struct_version >= 3)
SSLSocket_error("SSL_set_tlsext_host_name", NULL, net->socket, rc, opts->ssl_error_cb, opts->ssl_error_context);
else
SSLSocket_error("SSL_set_tlsext_host_name", NULL, net->socket, rc, NULL, NULL);
}
free(hostname_plus_null);
}
@ -661,7 +690,7 @@ int SSLSocket_setSocketForSSL(networkHandles* net, MQTTClient_SSLOptions* opts,
/*
* Return value: 1 - success, TCPSOCKET_INTERRUPTED - try again, anything else is failure
*/
int SSLSocket_connect(SSL* ssl, int sock, const char* hostname, int verify)
int SSLSocket_connect(SSL* ssl, int sock, const char* hostname, int verify, int (*cb)(const char *str, size_t len, void *u), void* u)
{
int rc = 0;
@ -671,7 +700,7 @@ int SSLSocket_connect(SSL* ssl, int sock, const char* hostname, int verify)
if (rc != 1)
{
int error;
error = SSLSocket_error("SSL_connect", ssl, sock, rc);
error = SSLSocket_error("SSL_connect", ssl, sock, rc, cb, u);
if (error == SSL_FATAL)
rc = error;
if (error == SSL_ERROR_WANT_READ || error == SSL_ERROR_WANT_WRITE)
@ -725,7 +754,7 @@ int SSLSocket_getch(SSL* ssl, int socket, char* c)
if ((rc = SSL_read(ssl, c, (size_t)1)) < 0)
{
int err = SSLSocket_error("SSL_read - getch", ssl, socket, rc);
int err = SSLSocket_error("SSL_read - getch", ssl, socket, rc, NULL, NULL);
if (err == SSL_ERROR_WANT_READ || err == SSL_ERROR_WANT_WRITE)
{
rc = TCPSOCKET_INTERRUPTED;
@ -770,7 +799,7 @@ char *SSLSocket_getdata(SSL* ssl, int socket, size_t bytes, size_t* actual_len)
if ((rc = SSL_read(ssl, buf + (*actual_len), (int)(bytes - (*actual_len)))) < 0)
{
rc = SSLSocket_error("SSL_read - getdata", ssl, socket, rc);
rc = SSLSocket_error("SSL_read - getdata", ssl, socket, rc, NULL, NULL);
if (rc != SSL_ERROR_WANT_READ && rc != SSL_ERROR_WANT_WRITE)
{
buf = NULL;
@ -865,7 +894,7 @@ int SSLSocket_putdatas(SSL* ssl, int socket, char* buf0, size_t buf0len, int cou
rc = TCPSOCKET_COMPLETE;
else
{
sslerror = SSLSocket_error("SSL_write", ssl, socket, rc);
sslerror = SSLSocket_error("SSL_write", ssl, socket, rc, NULL, NULL);
if (sslerror == SSL_ERROR_WANT_WRITE)
{
@ -948,7 +977,7 @@ int SSLSocket_continueWrite(pending_writes* pw)
}
else
{
int sslerror = SSLSocket_error("SSL_write", pw->ssl, pw->socket, rc);
int sslerror = SSLSocket_error("SSL_write", pw->ssl, pw->socket, rc, NULL, NULL);
if (sslerror == SSL_ERROR_WANT_WRITE)
rc = 0; /* indicate we haven't finished writing the payload yet */
}

View File

@ -44,7 +44,7 @@ char *SSLSocket_getdata(SSL* ssl, int socket, size_t bytes, size_t* actual_len);
int SSLSocket_close(networkHandles* net);
int SSLSocket_putdatas(SSL* ssl, int socket, char* buf0, size_t buf0len, int count, char** buffers, size_t* buflens, int* frees);
int SSLSocket_connect(SSL* ssl, int sock, const char* hostname, int verify);
int SSLSocket_connect(SSL* ssl, int sock, const char* hostname, int verify, int (*cb)(const char *str, size_t len, void *u), void* u);
int SSLSocket_getPendingRead(void);
int SSLSocket_continueWrite(pending_writes* pw);

View File

@ -211,6 +211,13 @@ void onPublish(void* context, MQTTAsync_successData* response)
}
static int onSSLError(const char *str, size_t len, void *context)
{
MQTTAsync client = (MQTTAsync)context;
return fprintf(stderr, "SSL error: %s\n", str);
}
void myconnect(MQTTAsync client)
{
MQTTAsync_connectOptions conn_opts = MQTTAsync_connectOptions_initializer;
@ -261,6 +268,8 @@ void myconnect(MQTTAsync client)
ssl_opts.privateKey = opts.key;
ssl_opts.privateKeyPassword = opts.keypass;
ssl_opts.enabledCipherSuites = opts.ciphers;
ssl_opts.ssl_error_cb = onSSLError;
ssl_opts.ssl_error_context = client;
conn_opts.ssl = &ssl_opts;
}